Skip to main content

tenferro_linalg/
traced.rs

1use std::sync::Arc;
2
3use num_complex::{Complex32, Complex64};
4use tenferro_runtime::extension::apply;
5use tenferro_runtime::{
6    CompareDir, DType, DotGeneralConfig, Error, ErrorPhase, Result, TracedTensor,
7};
8
9use crate::extension::{
10    validate_derivative_eps, EighOptions, LinalgExtensionOp, LinalgOp, QrOptions, SvdOptions,
11};
12use crate::rank_revealing_qr::validate_rank_revealing_qr_options;
13use crate::validation::{ensure_float_or_complex, validate_lstsq};
14use crate::{RankRevealingQrOptions, RankRevealingQrResult};
15
16/// Linear algebra extension methods for [`TracedTensor`].
17pub trait TracedTensorLinalgExt {
18    /// Build a traced SVD operation with default options.
19    ///
20    /// # Errors
21    ///
22    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
23    /// unsupported dtype, or `Error::Validation` for invalid graph metadata.
24    ///
25    /// # Deferred errors
26    ///
27    /// Backend numerical failures and concrete shape mismatches can be
28    /// reported as `Error::Extension` or `Error::Validation` during compile or
29    /// execution when symbolic inputs are bound.
30    fn svd(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor)>;
31
32    /// Build a traced SVD operation with explicit derivative and gauge options.
33    ///
34    /// # Errors
35    ///
36    /// Returns `Error::Validation::InvalidArgument` for a non-finite or
37    /// non-positive derivative epsilon, or `Error::Extension` for unsupported
38    /// dtype and graph registration failures.
39    ///
40    /// # Deferred errors
41    ///
42    /// Solver convergence and symbolic shape checks may be reported during
43    /// compile or execution.
44    fn svd_with_options(
45        &self,
46        options: SvdOptions,
47    ) -> Result<(TracedTensor, TracedTensor, TracedTensor)>;
48
49    /// Build a traced full-matrices SVD operation returning square `U (m x m)`
50    /// and `Vh (n x n)`, whose trailing `n - rank` rows span the input's right
51    /// nullspace.
52    ///
53    /// # Errors
54    ///
55    /// Returns `Error::Validation` when the input is not a batched matrix
56    /// (rank `>= 2`), `Error::Extension` with `ErrorKind::Unsupported` for
57    /// integer or boolean dtypes, or `Error::Extension` for graph
58    /// registration failures.
59    ///
60    /// # Deferred errors
61    ///
62    /// The active backend returns `Error::Extension` with
63    /// `ErrorKind::Unsupported` at execution if it does not implement
64    /// full-matrices SVD; both CPU providers and the CUDA backend implement
65    /// it. Automatic differentiation is intentionally unsupported for the full
66    /// variant (see the linalg AD support manifest) and surfaces a typed AD
67    /// error rather than a silent thin-SVD fallback.
68    fn svd_full(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor)>;
69
70    /// Build a traced QR operation.
71    ///
72    /// # Errors
73    ///
74    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
75    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
76    ///
77    /// # Deferred errors
78    ///
79    /// Concrete shape validation and backend QR failures may be reported at
80    /// compile or execution time for symbolic inputs.
81    fn qr(&self) -> Result<(TracedTensor, TracedTensor)>;
82
83    /// Build opaque compact Householder QR state.
84    ///
85    /// # Errors
86    ///
87    /// Returns `Error::Validation` for known invalid graph metadata or
88    /// `Error::Extension` for an unsupported operation.
89    ///
90    /// # Deferred errors
91    ///
92    /// Symbolic shape and backend provider checks may fail at compile or execution.
93    ///
94    /// # Examples
95    ///
96    /// ```rust
97    /// use tenferro_linalg::TracedTensorLinalgExt;
98    /// use tenferro_runtime::TracedTensor;
99    /// let a = TracedTensor::from_vec_col_major(vec![2, 1], vec![1.0_f64, 2.0])?;
100    /// let qr = a.householder_qr()?;
101    /// assert!(format!("{qr:?}").starts_with("HouseholderQr"));
102    /// # Ok::<(), tenferro_runtime::Error>(())
103    /// ```
104    fn householder_qr(&self) -> Result<crate::HouseholderQr<TracedTensor>>;
105
106    /// Build a traced QR operation with explicit gauge options.
107    ///
108    /// # Errors
109    ///
110    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
111    /// unsupported dtype, or `Error::Validation` for invalid graph metadata.
112    ///
113    /// # Deferred errors
114    ///
115    /// Symbolic shape checks and backend QR failures can be deferred to compile
116    /// or execution.
117    fn qr_with_options(&self, options: QrOptions) -> Result<(TracedTensor, TracedTensor)>;
118
119    /// Build fixed-arity traced column-pivoted rank-revealing QR.
120    ///
121    /// # Errors
122    /// Returns graph-build validation errors for rank, dtype, or invalid
123    /// tolerances, and extension registration failures.
124    ///
125    /// # Deferred errors
126    /// Symbolic shape checks, non-finite numerical failures, and unsupported
127    /// backend execution are reported during compile or execution.
128    ///
129    /// # Examples
130    ///
131    /// ```rust
132    /// use tenferro_linalg::{RankRevealingQrOptions, TracedTensorLinalgExt};
133    /// use tenferro_runtime::TracedTensor;
134    /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64, 0.0, 0.0, 1.0, 1.0, 1.0])?;
135    /// let result = a.rank_revealing_qr(RankRevealingQrOptions::default())?;
136    /// assert_eq!(result.q.rank, 2);
137    /// assert_eq!(result.column_permutation.rank, 1);
138    /// assert_eq!(result.rank.rank, 0);
139    /// # Ok::<(), tenferro_runtime::Error>(())
140    /// ```
141    fn rank_revealing_qr(
142        &self,
143        options: RankRevealingQrOptions,
144    ) -> Result<RankRevealingQrResult<TracedTensor>>;
145
146    /// Build a traced Hermitian eigendecomposition operation.
147    ///
148    /// # Errors
149    ///
150    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
151    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
152    ///
153    /// # Deferred errors
154    ///
155    /// Concrete square-shape validation and solver failures may be reported at
156    /// compile or execution time.
157    fn eigh(&self) -> Result<(TracedTensor, TracedTensor)>;
158
159    /// Build a traced Hermitian eigendecomposition with explicit options.
160    ///
161    /// # Errors
162    ///
163    /// Returns `Error::Validation::InvalidArgument` for an invalid derivative
164    /// epsilon, or `Error::Extension` for unsupported dtype and registration
165    /// failures.
166    ///
167    /// # Deferred errors
168    ///
169    /// Symbolic square-shape checks and numerical eigensolver failures may be
170    /// reported during compile or execution.
171    fn eigh_with_options(&self, options: EighOptions) -> Result<(TracedTensor, TracedTensor)>;
172
173    /// Build a traced Cholesky factorization operation.
174    ///
175    /// # Errors
176    ///
177    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
178    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
179    ///
180    /// # Deferred errors
181    ///
182    /// Non-square or non-positive-definite concrete inputs can produce
183    /// validation or numerical extension errors during compile or execution.
184    fn cholesky(&self) -> Result<TracedTensor>;
185
186    /// Build a traced LU factorization operation.
187    ///
188    /// # Errors
189    ///
190    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
191    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
192    ///
193    /// # Deferred errors
194    ///
195    /// Concrete shape checks and backend factorization failures may be
196    /// reported during compile or execution.
197    fn lu(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor, TracedTensor)>;
198
199    /// Build a traced complete-pivot LU factorization operation.
200    ///
201    /// # Errors
202    ///
203    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
204    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
205    ///
206    /// # Deferred errors
207    ///
208    /// Concrete square-shape checks and backend factorization failures may be
209    /// reported during compile or execution.
210    fn full_piv_lu(
211        &self,
212    ) -> Result<(
213        TracedTensor,
214        TracedTensor,
215        TracedTensor,
216        TracedTensor,
217        TracedTensor,
218    )>;
219    /// Build a traced general eigendecomposition operation.
220    ///
221    /// # Errors
222    ///
223    /// Returns `Error::Extension` with `ErrorKind::Unsupported` for an
224    /// unsupported dtype or `Error::Validation` for invalid graph metadata.
225    ///
226    /// # Deferred errors
227    ///
228    /// Concrete shape validation and numerical eigensolver failures may be
229    /// reported during compile or execution.
230    fn eig(&self) -> Result<(TracedTensor, TracedTensor)>;
231
232    /// Build a traced linear solve operation.
233    ///
234    /// # Errors
235    ///
236    /// Returns `Error::Validation` for incompatible coefficient/rhs metadata
237    /// and `Error::Extension` for unsupported dtype or registration failures.
238    ///
239    /// # Deferred errors
240    ///
241    /// Singular systems and concrete shape mismatches are reported as
242    /// numerical or validation errors during compile or execution.
243    fn solve(&self, b: &TracedTensor) -> Result<TracedTensor>;
244
245    /// Build a traced least-squares solve `argmin_x ||A x - b||_2` for a tall
246    /// or square, full-column-rank `A`, via the thin QR factorization.
247    ///
248    /// # Errors
249    ///
250    /// Returns `Error::Validation` for an invalid rank (`A` or `b` not a
251    /// batched matrix, rank `< 2`), a symbolic shape, a wide/underdetermined
252    /// `A` (`rows < cols`), or an unsupported dtype (not floating-point or
253    /// complex).
254    ///
255    /// # Deferred errors
256    ///
257    /// Backend QR and triangular-solve failures and concrete shape mismatches
258    /// are reported during compile or execution. Rank-deficient `A` is not
259    /// detected: `R` is singular and the result is ill-defined, so callers must
260    /// ensure full column rank.
261    fn lstsq(&self, b: &TracedTensor) -> Result<TracedTensor>;
262
263    /// Build a traced complete-pivot LU solve operation.
264    ///
265    /// # Errors
266    ///
267    /// Returns `Error::Validation` for incompatible coefficient/rhs metadata
268    /// and `Error::Extension` for unsupported dtype or registration failures.
269    ///
270    /// # Deferred errors
271    ///
272    /// Singular systems and concrete shape mismatches may be reported during
273    /// compile or execution.
274    fn full_piv_lu_solve(&self, b: &TracedTensor) -> Result<TracedTensor>;
275
276    /// Build a traced triangular solve operation.
277    ///
278    /// # Errors
279    ///
280    /// Returns `Error::Validation` for incompatible coefficient/rhs shapes or
281    /// invalid solve flags, and `Error::Extension` for unsupported dtype.
282    ///
283    /// # Deferred errors
284    ///
285    /// Singular or zero-diagonal systems can fail numerically during compile or
286    /// execution after symbolic inputs are bound.
287    fn triangular_solve(
288        &self,
289        b: &TracedTensor,
290        left_side: bool,
291        lower: bool,
292        transpose_a: bool,
293        unit_diagonal: bool,
294    ) -> Result<TracedTensor>;
295    /// Build a traced sign/log-determinant operation.
296    ///
297    /// # Errors
298    ///
299    /// Returns `Error::Validation` for invalid matrix metadata or
300    /// `Error::Extension` for unsupported dtype and registration failures.
301    ///
302    /// # Deferred errors
303    ///
304    /// Concrete singularity and shape failures can be reported during compile
305    /// or execution.
306    fn slogdet(&self) -> Result<(TracedTensor, TracedTensor)>;
307
308    /// Build a traced determinant operation.
309    ///
310    /// # Errors
311    ///
312    /// Returns `Error::Validation` for invalid matrix metadata or
313    /// `Error::Extension` for unsupported dtype.
314    ///
315    /// # Deferred errors
316    ///
317    /// Concrete singularity and shape failures may be reported during compile
318    /// or execution.
319    fn det(&self) -> Result<TracedTensor>;
320
321    /// Build a traced matrix-inverse operation.
322    ///
323    /// # Errors
324    ///
325    /// Returns `Error::Validation` for incompatible rank/shape metadata or
326    /// `Error::Extension` for unsupported dtype.
327    ///
328    /// # Deferred errors
329    ///
330    /// Singular matrices produce a numerical error during compile or execution.
331    fn inv(&self) -> Result<TracedTensor>;
332
333    /// Build a traced Hermitian eigenvalue-only operation.
334    ///
335    /// # Errors
336    ///
337    /// Returns `Error::Validation` for non-square metadata or
338    /// `Error::Extension` for unsupported dtype.
339    ///
340    /// # Deferred errors
341    ///
342    /// Concrete square-shape and solver failures may be reported during compile
343    /// or execution.
344    fn eigvalsh(&self) -> Result<TracedTensor>;
345
346    /// Build a traced general eigenvalue-only operation.
347    ///
348    /// # Errors
349    ///
350    /// Returns `Error::Validation` for invalid matrix metadata or
351    /// `Error::Extension` for unsupported dtype.
352    ///
353    /// # Deferred errors
354    ///
355    /// Concrete shape and eigensolver failures may be reported during compile
356    /// or execution.
357    fn eigvals(&self) -> Result<TracedTensor>;
358
359    /// Build a traced pseudoinverse operation with the default tolerance.
360    ///
361    /// # Errors
362    ///
363    /// Returns `Error::Validation` for invalid rank/shape metadata or
364    /// `Error::Extension` for unsupported dtype.
365    ///
366    /// # Deferred errors
367    ///
368    /// SVD convergence and concrete shape failures may be reported during
369    /// compile or execution.
370    fn pinv(&self) -> Result<TracedTensor>;
371
372    /// Build a traced pseudoinverse with an explicit relative tolerance.
373    ///
374    /// # Errors
375    ///
376    /// Returns `Error::Validation::InvalidArgument` when `rtol` is non-finite
377    /// or negative, or `Error::Extension` for unsupported dtype.
378    ///
379    /// # Deferred errors
380    ///
381    /// SVD convergence and concrete shape failures may be reported during
382    /// compile or execution.
383    fn pinv_with_rtol(&self, rtol: f64) -> Result<TracedTensor>;
384
385    /// Build a traced vector/matrix norm operation.
386    ///
387    /// # Errors
388    ///
389    /// Requires a concrete shape immediately. Returns `Error::Validation` for
390    /// invalid or duplicate axes or an invalid norm order. Symbolic shapes
391    /// produce `Error::TensorRuntime` wrapping `ValidationError::InvalidArgument`
392    /// for `shape` during graph construction; unsupported dtypes produce
393    /// `Error::Extension`.
394    ///
395    /// # Deferred errors
396    ///
397    /// Backend numerical or runtime failures may occur during execution.
398    fn norm(&self, ord: Option<f64>, dim: Option<&[usize]>, keepdim: bool) -> Result<TracedTensor>;
399}
400
401impl TracedTensorLinalgExt for TracedTensor {
402    fn svd(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
403        svd(self)
404    }
405
406    fn svd_with_options(
407        &self,
408        options: SvdOptions,
409    ) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
410        svd_with_options(self, options)
411    }
412
413    fn svd_full(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
414        svd_full(self)
415    }
416
417    fn qr(&self) -> Result<(TracedTensor, TracedTensor)> {
418        qr(self)
419    }
420
421    fn householder_qr(&self) -> Result<crate::HouseholderQr<TracedTensor>> {
422        householder_qr(self)
423    }
424
425    fn qr_with_options(&self, options: QrOptions) -> Result<(TracedTensor, TracedTensor)> {
426        qr_with_options(self, options)
427    }
428
429    fn rank_revealing_qr(
430        &self,
431        options: RankRevealingQrOptions,
432    ) -> Result<RankRevealingQrResult<TracedTensor>> {
433        rank_revealing_qr(self, options)
434    }
435
436    fn eigh(&self) -> Result<(TracedTensor, TracedTensor)> {
437        eigh(self)
438    }
439
440    fn eigh_with_options(&self, options: EighOptions) -> Result<(TracedTensor, TracedTensor)> {
441        eigh_with_options(self, options)
442    }
443
444    fn cholesky(&self) -> Result<TracedTensor> {
445        cholesky(self)
446    }
447
448    fn lu(&self) -> Result<(TracedTensor, TracedTensor, TracedTensor, TracedTensor)> {
449        lu(self)
450    }
451
452    fn full_piv_lu(
453        &self,
454    ) -> Result<(
455        TracedTensor,
456        TracedTensor,
457        TracedTensor,
458        TracedTensor,
459        TracedTensor,
460    )> {
461        full_piv_lu(self)
462    }
463
464    fn eig(&self) -> Result<(TracedTensor, TracedTensor)> {
465        eig(self)
466    }
467
468    fn solve(&self, b: &TracedTensor) -> Result<TracedTensor> {
469        solve(self, b)
470    }
471
472    fn lstsq(&self, b: &TracedTensor) -> Result<TracedTensor> {
473        lstsq(self, b)
474    }
475
476    fn full_piv_lu_solve(&self, b: &TracedTensor) -> Result<TracedTensor> {
477        full_piv_lu_solve(self, b)
478    }
479
480    fn triangular_solve(
481        &self,
482        b: &TracedTensor,
483        left_side: bool,
484        lower: bool,
485        transpose_a: bool,
486        unit_diagonal: bool,
487    ) -> Result<TracedTensor> {
488        triangular_solve(self, b, left_side, lower, transpose_a, unit_diagonal)
489    }
490
491    fn slogdet(&self) -> Result<(TracedTensor, TracedTensor)> {
492        slogdet(self)
493    }
494
495    fn det(&self) -> Result<TracedTensor> {
496        det(self)
497    }
498
499    fn inv(&self) -> Result<TracedTensor> {
500        inv(self)
501    }
502
503    fn eigvalsh(&self) -> Result<TracedTensor> {
504        eigvalsh(self)
505    }
506
507    fn eigvals(&self) -> Result<TracedTensor> {
508        eigvals(self)
509    }
510
511    fn pinv(&self) -> Result<TracedTensor> {
512        pinv(self)
513    }
514
515    fn pinv_with_rtol(&self, rtol: f64) -> Result<TracedTensor> {
516        pinv_with_rtol(self, rtol)
517    }
518
519    fn norm(&self, ord: Option<f64>, dim: Option<&[usize]>, keepdim: bool) -> Result<TracedTensor> {
520        norm(self, ord, dim, keepdim)
521    }
522}
523
524/// Build a traced singular value decomposition op using default options.
525///
526/// # Examples
527///
528/// ```
529/// use tenferro_linalg::TracedTensorLinalgExt;
530/// use tenferro_runtime::TracedTensor;
531///
532/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0]).unwrap();
533/// let (u, s, vt) = a.svd().unwrap();
534/// assert_eq!(u.rank, 2);
535/// assert_eq!(s.rank, 1);
536/// assert_eq!(vt.rank, 2);
537/// ```
538///
539/// # Errors
540///
541/// Returns `Error::Validation` for a known invalid rank, matrix shape, or
542/// dtype, `Error::Extension` with an unsupported-dtype or non-convergence
543/// source when the registered linalg backend cannot construct the operation,
544/// and `Error::RuntimeState` when extension registration is unavailable.
545///
546/// # Deferred errors
547///
548/// A symbolic matrix or batch-shape mismatch is reported later as
549/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation` during compile or
550/// execution.
551pub fn svd(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
552    svd_with_options(a, SvdOptions::default())
553}
554
555/// Build a traced singular value decomposition op with explicit options.
556///
557/// `derivative_eps` regularizes decomposition derivative formulas. It is not a
558/// backend SVD solver tolerance.
559///
560/// # Examples
561///
562/// ```
563/// use tenferro_linalg::{SvdGauge, SvdOptions, TracedTensorLinalgExt};
564/// use tenferro_runtime::TracedTensor;
565///
566/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0]).unwrap();
567/// let options = SvdOptions::default()
568///     .gauge(SvdGauge::CanonicalPivot)
569///     .derivative_eps(1e-10);
570/// let (_u, s, _vt) = a.svd_with_options(options).unwrap();
571/// assert_eq!(s.rank, 1);
572/// ```
573///
574/// # Errors
575///
576/// Returns `Error::Validation` when `derivative_eps` is non-finite or
577/// non-positive, `Error::Extension` for an unsupported dtype or numerical
578/// non-convergence, and `Error::Internal` if the extension output contract is
579/// violated.
580///
581/// # Deferred errors
582///
583/// Symbolic rank or shape constraints are checked later and can produce
584/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
585pub fn svd_with_options(
586    a: &TracedTensor,
587    options: SvdOptions,
588) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
589    validate_derivative_eps("svd_with_options", options.derivative_eps)?;
590    ensure_float_or_complex("svd", a.dtype)?;
591    three_outputs(
592        apply(
593            Arc::new(LinalgExtensionOp::new(LinalgOp::Svd {
594                derivative_eps: options.derivative_eps,
595                gauge: options.gauge,
596                driver: options.driver,
597            })),
598            &[a],
599        )?,
600        "svd",
601    )
602}
603
604/// Build a traced full-matrices singular value decomposition op.
605///
606/// Unlike [`svd`], the returned factors are square: `U` is `m x m` and `Vh` is
607/// `n x n`, while `S` still holds `min(m, n)` singular values. The trailing
608/// `n - rank` rows of `Vh` span the right nullspace of the input, so this is
609/// the decomposition to use for kernel-basis extraction.
610///
611/// # Examples
612///
613/// ```
614/// use tenferro_linalg::TracedTensorLinalgExt;
615/// use tenferro_runtime::TracedTensor;
616///
617/// // A wide 1x2 system: the trailing row of the 2x2 Vh spans the nullspace.
618/// let a = TracedTensor::from_vec_col_major(vec![1, 2], vec![1.0_f64, 1.0]).unwrap();
619/// let (u, s, vh) = a.svd_full().unwrap();
620/// assert_eq!(u.rank, 2);
621/// assert_eq!(s.rank, 1);
622/// assert_eq!(vh.rank, 2);
623/// ```
624///
625/// # Errors
626///
627/// Returns `Error::Validation` when the input is not a batched matrix
628/// (rank `>= 2`) or `Error::Extension` with `ErrorKind::Unsupported` for
629/// integer or boolean dtypes; `Error::RuntimeState` when extension
630/// registration is unavailable.
631///
632/// # Deferred errors
633///
634/// The active backend returns `Error::Extension` with `ErrorKind::Unsupported`
635/// during execution if it does not implement full-matrices SVD; both CPU
636/// providers and the CUDA backend implement it. Automatic differentiation is
637/// intentionally unsupported for the full variant (see the linalg AD support
638/// manifest) and surfaces a typed AD error, not a silent thin-SVD fallback.
639pub fn svd_full(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
640    ensure_float_or_complex("svd_full", a.dtype)?;
641    three_outputs(
642        apply(Arc::new(LinalgExtensionOp::new(LinalgOp::SvdFull)), &[a])?,
643        "svd_full",
644    )
645}
646
647/// Build a traced QR decomposition op.
648///
649/// # Examples
650///
651/// ```
652/// use tenferro_linalg::TracedTensorLinalgExt;
653/// use tenferro_runtime::TracedTensor;
654///
655/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0]).unwrap();
656/// let (q, r) = a.qr().unwrap();
657/// assert_eq!(q.rank, 2);
658/// assert_eq!(r.rank, 2);
659/// ```
660///
661/// # Errors
662///
663/// Returns `Error::Validation` for a known invalid rank or matrix shape,
664/// `Error::Extension` for an unsupported dtype or numerical failure, and
665/// `Error::RuntimeState` when the linalg extension is not registered.
666///
667/// # Deferred errors
668///
669/// Unknown matrix or batch dimensions can fail later as
670/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
671pub fn qr(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor)> {
672    qr_with_options(a, QrOptions::default())
673}
674
675/// Build compact Householder QR state for a traced matrix.
676///
677/// # Errors
678///
679/// Returns `Error::Validation` for known invalid graph metadata or
680/// `Error::Extension` for an unsupported operation.
681///
682/// # Deferred errors
683///
684/// Symbolic shape constraints and backend provider failures may be reported
685/// during compile or execution.
686pub fn householder_qr(a: &TracedTensor) -> Result<crate::HouseholderQr<TracedTensor>> {
687    ensure_float_or_complex("householder_qr", a.dtype)?;
688    let mut outputs = apply(
689        Arc::new(LinalgExtensionOp::new(LinalgOp::HouseholderQrFactor)),
690        &[a],
691    )?
692    .into_iter();
693    match (outputs.next(), outputs.next(), outputs.next()) {
694        (Some(packed), Some(coeff), None) => Ok(
695            crate::HouseholderQr::<TracedTensor>::from_traced_outputs(packed, coeff),
696        ),
697        _ => Err(unexpected_output_count("householder_qr", 2)),
698    }
699}
700
701/// Build a traced QR decomposition op with explicit options.
702///
703/// `gauge` controls optional sign or phase post-processing.
704///
705/// # Examples
706///
707/// ```
708/// use tenferro_linalg::{QrGauge, QrOptions, TracedTensorLinalgExt};
709/// use tenferro_runtime::TracedTensor;
710///
711/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0]).unwrap();
712/// let (q, r) = a.qr_with_options(QrOptions::default().gauge(QrGauge::PositiveDiagonal)).unwrap();
713/// assert_eq!(q.rank, 2);
714/// assert_eq!(r.rank, 2);
715/// ```
716///
717/// # Errors
718///
719/// Returns `Error::Validation` for a known invalid rank or matrix shape,
720/// `Error::Extension` for an unsupported dtype or numerical failure, and
721/// `Error::Internal` if the extension output contract is violated.
722///
723/// # Deferred errors
724///
725/// Symbolic matrix or batch constraints are checked later and can produce
726/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
727pub fn qr_with_options(
728    a: &TracedTensor,
729    options: QrOptions,
730) -> Result<(TracedTensor, TracedTensor)> {
731    ensure_float_or_complex("qr", a.dtype)?;
732    two_outputs(
733        apply(
734            Arc::new(LinalgExtensionOp::new(LinalgOp::Qr {
735                gauge: options.gauge,
736            })),
737            &[a],
738        )?,
739        "qr",
740    )
741}
742
743/// Build a traced column-pivoted rank-revealing QR operation.
744///
745/// # Errors
746/// Returns graph-build validation errors for invalid rank, dtype, or
747/// tolerances, plus extension registration failures.
748///
749/// # Deferred errors
750/// Symbolic shape checks, non-finite numerical failures, and unsupported
751/// backend execution are reported during compile or execution.
752///
753/// # Examples
754///
755/// ```rust
756/// use tenferro_linalg::{RankRevealingQrOptions, TracedTensorLinalgExt};
757/// use tenferro_runtime::TracedTensor;
758/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
759/// let result = a.rank_revealing_qr(RankRevealingQrOptions::default())?;
760/// assert_eq!(result.column_permutation.rank, 1);
761/// assert_eq!(result.rank.rank, 0);
762/// # Ok::<(), tenferro_runtime::Error>(())
763/// ```
764pub fn rank_revealing_qr(
765    a: &TracedTensor,
766    options: RankRevealingQrOptions,
767) -> Result<RankRevealingQrResult<TracedTensor>> {
768    validate_rank_revealing_qr_options("rank_revealing_qr", options)?;
769    ensure_float_or_complex("rank_revealing_qr", a.dtype)?;
770    let (q, r, column_permutation, rank) = four_outputs(
771        apply(
772            Arc::new(LinalgExtensionOp::new(LinalgOp::RankRevealingQr {
773                gauge: options.gauge,
774                rtol: options.rtol,
775                atol: options.atol,
776            })),
777            &[a],
778        )?,
779        "rank_revealing_qr",
780    )?;
781    Ok(RankRevealingQrResult {
782        q,
783        r,
784        column_permutation,
785        rank,
786    })
787}
788
789/// Build a traced Hermitian eigenvalue decomposition op using default options.
790///
791/// # Examples
792///
793/// ```
794/// use tenferro_linalg::TracedTensorLinalgExt;
795/// use tenferro_runtime::TracedTensor;
796///
797/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
798/// let (values, vectors) = a.eigh().unwrap();
799/// assert_eq!(values.rank, 1);
800/// assert_eq!(vectors.rank, 2);
801/// ```
802///
803/// # Errors
804///
805/// Returns `Error::Validation` for a known non-square or invalid-rank input,
806/// `Error::Extension` for an unsupported dtype or eigensolver
807/// non-convergence, and `Error::RuntimeState` when the extension is not
808/// registered.
809///
810/// # Deferred errors
811///
812/// Symbolic square-shape constraints can fail later as
813/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
814pub fn eigh(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor)> {
815    eigh_with_options(a, EighOptions::default())
816}
817
818/// Build a traced Hermitian eigenvalue decomposition op with explicit options.
819///
820/// `derivative_eps` regularizes derivative formulas for repeated or nearly
821/// repeated eigenvalues. It is not a backend eigensolver tolerance.
822///
823/// # Examples
824///
825/// ```
826/// use tenferro_linalg::{EighGauge, EighOptions, TracedTensorLinalgExt};
827/// use tenferro_runtime::TracedTensor;
828///
829/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
830/// let (values, _vectors) = a
831///     .eigh_with_options(
832///         EighOptions::default()
833///             .gauge(EighGauge::CanonicalPivot)
834///             .derivative_eps(1e-10),
835///     )
836///     .unwrap();
837/// assert_eq!(values.rank, 1);
838/// ```
839///
840/// # Errors
841///
842/// Returns `Error::Validation` for a known non-square or invalid-rank input,
843/// or for non-finite/non-positive `derivative_eps`; `Error::Extension` for an
844/// unsupported dtype or eigensolver non-convergence; and `Error::Internal` for
845/// an output-count contract violation.
846///
847/// # Deferred errors
848///
849/// Symbolic square-shape constraints can fail later as
850/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
851pub fn eigh_with_options(
852    a: &TracedTensor,
853    options: EighOptions,
854) -> Result<(TracedTensor, TracedTensor)> {
855    validate_derivative_eps("eigh_with_options", options.derivative_eps)?;
856    ensure_float_or_complex("eigh", a.dtype)?;
857    two_outputs(
858        apply(
859            Arc::new(LinalgExtensionOp::new(LinalgOp::Eigh {
860                derivative_eps: options.derivative_eps,
861                gauge: options.gauge,
862                driver: options.driver,
863            })),
864            &[a],
865        )?,
866        "eigh",
867    )
868}
869
870/// Build a traced Cholesky decomposition op.
871///
872/// # Examples
873///
874/// ```
875/// use tenferro_linalg::TracedTensorLinalgExt;
876/// use tenferro_runtime::TracedTensor;
877///
878/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![4.0_f64, 2.0, 2.0, 3.0]).unwrap();
879/// let factor = a.cholesky().unwrap();
880/// assert_eq!(factor.rank, 2);
881/// ```
882///
883/// # Errors
884///
885/// Returns `Error::Validation` for a known non-square or invalid-rank input,
886/// `Error::Extension` for an unsupported dtype or a non-positive-definite
887/// matrix, and `Error::RuntimeState` when the extension is not registered.
888///
889/// # Deferred errors
890///
891/// Symbolic square-shape constraints can fail later as
892/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
893pub fn cholesky(a: &TracedTensor) -> Result<TracedTensor> {
894    ensure_float_or_complex("cholesky", a.dtype)?;
895    one_output(
896        apply(Arc::new(LinalgExtensionOp::new(LinalgOp::Cholesky)), &[a])?,
897        "cholesky",
898    )
899}
900
901/// Build a traced LU decomposition op.
902///
903/// # Examples
904///
905/// ```
906/// use tenferro_linalg::TracedTensorLinalgExt;
907/// use tenferro_runtime::TracedTensor;
908///
909/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0]).unwrap();
910/// let (p, l, u, parity) = a.lu().unwrap();
911/// assert_eq!(p.rank, 2);
912/// assert_eq!(l.rank, 2);
913/// assert_eq!(u.rank, 2);
914/// assert_eq!(parity.rank, 0);
915/// ```
916///
917/// # Errors
918///
919/// Returns `Error::Validation` for a known invalid rank or matrix shape,
920/// `Error::Extension` for an unsupported dtype or singular numerical result,
921/// and `Error::RuntimeState` when the extension is not registered.
922///
923/// # Deferred errors
924///
925/// Symbolic square-shape constraints can fail later as
926/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
927pub fn lu(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor, TracedTensor, TracedTensor)> {
928    four_outputs(
929        apply(Arc::new(LinalgExtensionOp::new(LinalgOp::Lu)), &[a])?,
930        "lu",
931    )
932}
933
934/// Build a traced full-pivot LU decomposition op.
935///
936/// Returns `(P, L, U, Q, parity)` with reconstruction convention
937/// `A = P^T * L * U * Q`, equivalently `P * A * Q^T = L * U`. `parity` is a
938/// scalar real tensor containing `+1` or `-1`: `F32` for `F32`/`C32` inputs and
939/// `F64` for `F64`/`C64` inputs.
940///
941/// # Examples
942///
943/// ```
944/// use tenferro_linalg::TracedTensorLinalgExt;
945/// use tenferro_runtime::TracedTensor;
946///
947/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0]).unwrap();
948/// let (p, l, u, q, parity) = a.full_piv_lu().unwrap();
949/// assert_eq!(p.rank, 2);
950/// assert_eq!(l.rank, 2);
951/// assert_eq!(u.rank, 2);
952/// assert_eq!(q.rank, 2);
953/// assert_eq!(parity.rank, 0);
954/// ```
955///
956/// # Errors
957///
958/// Returns `Error::Validation` for a known invalid rank or matrix shape,
959/// `Error::Extension` for an unsupported dtype or singular numerical result,
960/// and `Error::Internal` for an output-count contract violation.
961///
962/// # Deferred errors
963///
964/// Symbolic square-shape constraints can fail later as
965/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
966pub fn full_piv_lu(
967    a: &TracedTensor,
968) -> Result<(
969    TracedTensor,
970    TracedTensor,
971    TracedTensor,
972    TracedTensor,
973    TracedTensor,
974)> {
975    five_outputs(
976        apply(Arc::new(LinalgExtensionOp::new(LinalgOp::FullPivLu)), &[a])?,
977        "full_piv_lu",
978    )
979}
980
981/// Build a traced general eigendecomposition op.
982///
983/// # Examples
984///
985/// ```
986/// use tenferro_linalg::TracedTensorLinalgExt;
987/// use tenferro_runtime::TracedTensor;
988///
989/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0]).unwrap();
990/// let (values, vectors) = a.eig().unwrap();
991/// assert_eq!(values.rank, 1);
992/// assert_eq!(vectors.rank, 2);
993/// ```
994///
995/// # Errors
996///
997/// Returns `Error::Validation` for a known non-square or invalid-rank input,
998/// `Error::Extension` for an unsupported dtype or eigensolver
999/// non-convergence, and `Error::RuntimeState` when the extension is not
1000/// registered.
1001///
1002/// # Deferred errors
1003///
1004/// Symbolic square-shape constraints can fail later as
1005/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
1006pub fn eig(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor)> {
1007    two_outputs(
1008        apply(
1009            Arc::new(LinalgExtensionOp::new(LinalgOp::Eig {
1010                input_dtype: a.dtype,
1011            })),
1012            &[a],
1013        )?,
1014        "eig",
1015    )
1016}
1017
1018/// Build a traced linear solve op.
1019///
1020/// # Examples
1021///
1022/// ```
1023/// use tenferro_linalg::TracedTensorLinalgExt;
1024/// use tenferro_runtime::TracedTensor;
1025///
1026/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1027/// let b = TracedTensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 9.0]).unwrap();
1028/// let x = a.solve(&b).unwrap();
1029/// assert_eq!(x.rank, 2);
1030/// ```
1031///
1032/// # Errors
1033///
1034/// Returns `Error::Validation` for known incompatible matrix, batch, or dtype
1035/// metadata, `Error::Extension` for an unsupported dtype or singular system,
1036/// and `Error::RuntimeState` when the extension is not registered.
1037///
1038/// # Deferred errors
1039///
1040/// Symbolic matrix and batch constraints can fail later as
1041/// `ShapeConstraintViolation`, `ShapeConstraintEvaluation`, or
1042/// `ShapeExpressionEvaluation`.
1043pub fn solve(a: &TracedTensor, b: &TracedTensor) -> Result<TracedTensor> {
1044    // One fused op factors and solves in a single backend call and saves
1045    // (x, packed LU, pivots), so AD reuses the factors for the adjoint solve.
1046    let mut outputs = apply(
1047        Arc::new(LinalgExtensionOp::new(LinalgOp::LuFactorSolve)),
1048        &[a, b],
1049    )?
1050    .into_iter();
1051    match (
1052        outputs.next(),
1053        outputs.next(),
1054        outputs.next(),
1055        outputs.next(),
1056    ) {
1057        (Some(x), Some(_packed_lu), Some(_pivots), None) => Ok(x),
1058        _ => Err(unexpected_output_count("lu_factor_solve", 3)),
1059    }
1060}
1061
1062/// Build a traced least-squares solve `argmin_x ||A x - b||_2` for a tall or
1063/// square, full-column-rank `A`.
1064///
1065/// The solution is computed through the thin QR factorization `A = Q R`: since
1066/// `R` is nonsingular for full column rank, `x = R^{-1} (Qá´´ b)`. This composes
1067/// existing traced decomposition ops (`qr`, `dot_general`, `triangular_solve`),
1068/// so, unlike the value-only [`svd_full`], it participates in autodiff through
1069/// its component rules.
1070///
1071/// # Examples
1072///
1073/// ```
1074/// use tenferro_linalg::TracedTensorLinalgExt;
1075/// use tenferro_runtime::TracedTensor;
1076///
1077/// // Overdetermined 3x2 system.
1078/// let a = TracedTensor::from_vec_col_major(
1079///     vec![3, 2],
1080///     vec![1.0_f64, 1.0, 1.0, 0.0, 1.0, 2.0],
1081/// )
1082/// .unwrap();
1083/// let b = TracedTensor::from_vec_col_major(vec![3, 1], vec![1.0_f64, 2.0, 2.0]).unwrap();
1084/// let x = a.lstsq(&b).unwrap();
1085/// assert_eq!(x.rank, 2);
1086/// ```
1087///
1088/// # Errors
1089///
1090/// Returns `Error::Validation` when `A` or `b` is not a batched matrix
1091/// (rank `>= 2`), when `A` has a symbolic shape, when `A` is wide
1092/// (`rows < cols`, underdetermined), or when the dtype is not floating-point or
1093/// complex. Rank-deficient `A` is not detected here: `R` is singular and the
1094/// triangular solve yields a non-finite or ill-defined result, so callers must
1095/// ensure full column rank.
1096///
1097/// # Deferred errors
1098///
1099/// Backend QR and triangular-solve failures and concrete shape mismatches are
1100/// reported during compile or execution.
1101pub fn lstsq(a: &TracedTensor, b: &TracedTensor) -> Result<TracedTensor> {
1102    validate_lstsq(
1103        "lstsq",
1104        a.dtype,
1105        a.rank,
1106        b.rank,
1107        || {
1108            let shape = require_concrete_shape("lstsq", a)?;
1109            Ok((shape[0], shape[1]))
1110        },
1111        |message| {
1112            Error::TensorRuntime(tenferro_tensor::Error::invalid_argument(
1113                "lstsq", "shape", message,
1114            ))
1115        },
1116    )?;
1117    let (q, r) = qr(a)?;
1118    let qh = q.conj()?.transpose(&matrix_transpose_perm(q.rank))?;
1119    let qh_b = matmul_preserve_trailing_batch(&qh, b)?;
1120    triangular_solve(&r, &qh_b, true, false, false, false)
1121}
1122
1123/// Build a traced full-pivot LU solve op.
1124///
1125/// # Examples
1126///
1127/// ```
1128/// use tenferro_linalg::TracedTensorLinalgExt;
1129/// use tenferro_runtime::TracedTensor;
1130///
1131/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1132/// let b = TracedTensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 9.0]).unwrap();
1133/// let x = a.full_piv_lu_solve(&b).unwrap();
1134/// assert_eq!(x.rank, 2);
1135/// ```
1136///
1137/// # Errors
1138///
1139/// Returns `Error::Validation` for known incompatible matrix, batch, or dtype
1140/// metadata, `Error::Extension` for an unsupported dtype or singular system,
1141/// and `Error::RuntimeState` when the extension is not registered.
1142///
1143/// # Deferred errors
1144///
1145/// Symbolic matrix and batch constraints can fail later as
1146/// `ShapeConstraintViolation`, `ShapeConstraintEvaluation`, or
1147/// `ShapeExpressionEvaluation`.
1148pub fn full_piv_lu_solve(a: &TracedTensor, b: &TracedTensor) -> Result<TracedTensor> {
1149    one_output(
1150        apply(
1151            Arc::new(LinalgExtensionOp::new(LinalgOp::FullPivLuSolve {
1152                transpose_a: false,
1153            })),
1154            &[a, b],
1155        )?,
1156        "full_piv_lu_solve",
1157    )
1158}
1159
1160/// Build a traced triangular solve op.
1161///
1162/// # Examples
1163///
1164/// ```
1165/// use tenferro_linalg::TracedTensorLinalgExt;
1166/// use tenferro_runtime::TracedTensor;
1167///
1168/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 1.0, 3.0]).unwrap();
1169/// let b = TracedTensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 9.0]).unwrap();
1170/// let x = a.triangular_solve(&b, true, true, false, false).unwrap();
1171/// assert_eq!(x.rank, 2);
1172/// ```
1173///
1174/// # Errors
1175///
1176/// Returns `Error::Validation` for incompatible matrix, batch, or dtype
1177/// metadata, `Error::Extension` for an unsupported dtype or singular system,
1178/// and `Error::RuntimeState` when the extension is not registered.
1179///
1180/// # Deferred errors
1181///
1182/// Symbolic matrix and batch constraints can fail later as
1183/// `ShapeConstraintViolation`, `ShapeConstraintEvaluation`, or
1184/// `ShapeExpressionEvaluation`.
1185pub fn triangular_solve(
1186    a: &TracedTensor,
1187    b: &TracedTensor,
1188    left_side: bool,
1189    lower: bool,
1190    transpose_a: bool,
1191    unit_diagonal: bool,
1192) -> Result<TracedTensor> {
1193    one_output(
1194        apply(
1195            Arc::new(LinalgExtensionOp::new(LinalgOp::TriangularSolve {
1196                left_side,
1197                lower,
1198                transpose_a,
1199                unit_diagonal,
1200            })),
1201            &[a, b],
1202        )?,
1203        "triangular_solve",
1204    )
1205}
1206
1207/// Build traced sign and log-absolute-determinant ops.
1208///
1209/// # Examples
1210///
1211/// ```
1212/// use tenferro_linalg::TracedTensorLinalgExt;
1213/// use tenferro_runtime::TracedTensor;
1214///
1215/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1216/// let (sign, logabsdet) = a.slogdet().unwrap();
1217/// assert_eq!(sign.rank, 0);
1218/// assert_eq!(logabsdet.rank, 0);
1219/// ```
1220///
1221/// # Errors
1222///
1223/// Returns `Error::Validation` for a known non-square or invalid-rank input,
1224/// `Error::Extension` for an unsupported dtype or singular factorization, and
1225/// `Error::Internal` if the factorization output contract is violated.
1226///
1227/// # Deferred errors
1228///
1229/// Symbolic square-shape constraints can fail later as
1230/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
1231pub fn slogdet(a: &TracedTensor) -> Result<(TracedTensor, TracedTensor)> {
1232    if let Some(empty) = slogdet_empty_square(a)? {
1233        return Ok(empty);
1234    }
1235    let mut factor_outputs =
1236        apply(Arc::new(LinalgExtensionOp::new(LinalgOp::LuFactor)), &[a])?.into_iter();
1237    let (packed_lu, parity) = match (
1238        factor_outputs.next(),
1239        factor_outputs.next(),
1240        factor_outputs.next(),
1241        factor_outputs.next(),
1242    ) {
1243        (Some(packed_lu), Some(_pivots), Some(parity), None) => (packed_lu, parity),
1244        _ => return Err(unexpected_output_count("lu_factor", 3)),
1245    };
1246    let mut sign_outputs = apply(
1247        Arc::new(LinalgExtensionOp::new(LinalgOp::SignDetFromLuFactor)),
1248        &[a, &packed_lu, &parity],
1249    )?
1250    .into_iter();
1251    let sign = match (sign_outputs.next(), sign_outputs.next()) {
1252        (Some(sign), None) => sign,
1253        _ => return Err(unexpected_output_count("signdet_from_lu_factor", 1)),
1254    };
1255    let mut logabsdet_outputs = apply(
1256        Arc::new(LinalgExtensionOp::new(LinalgOp::LogAbsDetFromLuFactor)),
1257        &[a, &packed_lu],
1258    )?
1259    .into_iter();
1260    let logabsdet = match (logabsdet_outputs.next(), logabsdet_outputs.next()) {
1261        (Some(logabsdet), None) => logabsdet,
1262        _ => return Err(unexpected_output_count("logabsdet_from_lu_factor", 1)),
1263    };
1264    Ok((sign, logabsdet))
1265}
1266
1267/// Build a traced determinant op.
1268///
1269/// The value contract follows JAX: `det = sign * exp(logabsdet)` from the same
1270/// factorization that [`slogdet`] uses. This avoids intermediate product overflow
1271/// (for example `diag(1e200,1e200,1e-200,1e-200)` gives `1`), while retaining
1272/// ordinary floating-point roundoff and final exponential overflow/underflow.
1273///
1274/// # Examples
1275///
1276/// ```
1277/// use tenferro_linalg::TracedTensorLinalgExt;
1278/// use tenferro_runtime::TracedTensor;
1279///
1280/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1281/// let determinant = a.det().unwrap();
1282/// assert_eq!(determinant.rank, 0);
1283/// ```
1284///
1285/// # Errors
1286///
1287/// Returns the same `Error::Validation`, `Error::Extension`, and
1288/// `Error::RuntimeState` failures as [`slogdet`], including a singular
1289/// factorization and an invalid matrix shape.
1290///
1291/// # Deferred errors
1292///
1293/// Symbolic shape checks can later produce `ShapeConstraintViolation`,
1294/// `ShapeConstraintEvaluation`, or `ShapeExpressionEvaluation`.
1295pub fn det(a: &TracedTensor) -> Result<TracedTensor> {
1296    // JAX's recipe: `sign, logdet = slogdet(a); return sign * exp(logdet)`
1297    // (`_det` in `jax/_src/numpy/linalg.py`). Multiplying the LU diagonal
1298    // instead overflows before the magnitudes cancel when the determinant
1299    // spans an extreme dynamic range, so the two recipes disagree there.
1300    let (sign, logabsdet) = slogdet(a)?;
1301    &sign * &logabsdet.exp()?
1302}
1303
1304fn slogdet_empty_square(a: &TracedTensor) -> Result<Option<(TracedTensor, TracedTensor)>> {
1305    let Some(shape) = a.try_concrete_shape() else {
1306        return Ok(None);
1307    };
1308    if shape.len() < 2 || shape[0] != 0 || shape[1] != 0 {
1309        return Ok(None);
1310    }
1311    let batch_shape = shape[2..].to_vec();
1312    Ok(Some((
1313        filled_real(a.dtype, batch_shape.clone(), 1.0)?,
1314        filled_real(real_values_dtype(a.dtype), batch_shape, 0.0)?,
1315    )))
1316}
1317
1318/// Build a traced matrix inverse op.
1319///
1320/// # Examples
1321///
1322/// ```
1323/// use tenferro_linalg::TracedTensorLinalgExt;
1324/// use tenferro_runtime::TracedTensor;
1325///
1326/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1327/// let inverse = a.inv().unwrap();
1328/// assert_eq!(inverse.rank, 2);
1329/// ```
1330///
1331/// # Errors
1332///
1333/// Returns `Error::Validation` when the input is not at least rank two or is
1334/// not square, `Error::Extension` for an unsupported dtype or singular system,
1335/// and `Error::RuntimeState` when the extension is not registered.
1336///
1337/// # Deferred errors
1338///
1339/// A symbolic shape that cannot provide the identity size fails later as
1340/// `ShapeConstraintEvaluation` or `ShapeExpressionEvaluation`.
1341pub fn inv(a: &TracedTensor) -> Result<TracedTensor> {
1342    ensure_min_rank("inv", a.rank, 2)?;
1343    let shape = require_concrete_shape("inv", a)?;
1344    let eye = eye_like(a, shape[0])?;
1345    solve(a, &eye)
1346}
1347
1348/// Build a traced Hermitian eigenvalue-only op.
1349///
1350/// # Examples
1351///
1352/// ```
1353/// use tenferro_linalg::TracedTensorLinalgExt;
1354/// use tenferro_runtime::TracedTensor;
1355///
1356/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 3.0]).unwrap();
1357/// let values = a.eigvalsh().unwrap();
1358/// assert_eq!(values.rank, 1);
1359/// ```
1360///
1361/// # Errors
1362///
1363/// Returns `Error::Validation` for a known non-square or invalid-rank input,
1364/// `Error::Extension` for an unsupported dtype or eigensolver
1365/// non-convergence, and `Error::RuntimeState` when the extension is not
1366/// registered.
1367///
1368/// # Deferred errors
1369///
1370/// Symbolic square-shape constraints can fail later as
1371/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
1372pub fn eigvalsh(a: &TracedTensor) -> Result<TracedTensor> {
1373    ensure_float_or_complex("eigvalsh", a.dtype)?;
1374    eigh_values(a)
1375}
1376
1377/// Build a traced general eigenvalue-only op.
1378///
1379/// # Examples
1380///
1381/// ```
1382/// use tenferro_linalg::TracedTensorLinalgExt;
1383/// use tenferro_runtime::TracedTensor;
1384///
1385/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0]).unwrap();
1386/// let values = a.eigvals().unwrap();
1387/// assert_eq!(values.rank, 1);
1388/// ```
1389///
1390/// # Errors
1391///
1392/// Returns `Error::Validation` for a known non-square or invalid-rank input,
1393/// `Error::Extension` for an unsupported dtype or eigensolver
1394/// non-convergence, and `Error::RuntimeState` when the extension is not
1395/// registered.
1396///
1397/// # Deferred errors
1398///
1399/// Symbolic square-shape constraints can fail later as
1400/// `ShapeConstraintViolation` or `ShapeConstraintEvaluation`.
1401pub fn eigvals(a: &TracedTensor) -> Result<TracedTensor> {
1402    eig_values(a)
1403}
1404
1405/// Build a traced Moore-Penrose pseudoinverse op.
1406///
1407/// Floating-point and complex inputs are supported. Integer and boolean inputs
1408/// return an unsupported-dtype error.
1409///
1410/// # Examples
1411///
1412/// ```
1413/// use tenferro_linalg::TracedTensorLinalgExt;
1414/// use tenferro_runtime::TracedTensor;
1415///
1416/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0]).unwrap();
1417/// let inverse = a.pinv().unwrap();
1418/// assert_eq!(inverse.rank, 2);
1419/// ```
1420///
1421/// # Errors
1422///
1423/// Returns `Error::Validation` for an invalid rank, shape, or negative/non-
1424/// finite `rtol`, `Error::Extension` for unsupported integer or boolean dtypes,
1425/// numerical non-convergence, or a backend failure, and `Error::RuntimeState`
1426/// when the extension is not registered.
1427///
1428/// # Deferred errors
1429///
1430/// Symbolic shapes are materialized by this helper; failures are reported as
1431/// `ShapeConstraintEvaluation` or `ShapeExpressionEvaluation`.
1432pub fn pinv(a: &TracedTensor) -> Result<TracedTensor> {
1433    ensure_float_or_complex("pinv", a.dtype)?;
1434    let shape = require_concrete_shape("pinv", a)?;
1435    let max_dim = match (shape.first(), shape.get(1)) {
1436        (Some(&m), Some(&n)) => m.max(n),
1437        (Some(&m), None) => m,
1438        _ => 0,
1439    };
1440    pinv_with_rtol(a, default_pinv_rtol(a.dtype, max_dim))
1441}
1442
1443/// Build a traced Moore-Penrose pseudoinverse op with an explicit relative tolerance.
1444///
1445/// Floating-point and complex inputs are supported. Integer and boolean inputs
1446/// return an unsupported-dtype error.
1447///
1448/// # Examples
1449///
1450/// ```
1451/// use tenferro_linalg::TracedTensorLinalgExt;
1452/// use tenferro_runtime::TracedTensor;
1453///
1454/// let a = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0]).unwrap();
1455/// let inverse = a.pinv_with_rtol(1e-12).unwrap();
1456/// assert_eq!(inverse.rank, 2);
1457/// ```
1458///
1459/// # Errors
1460///
1461/// Returns `Error::Validation` for an invalid rank, shape, or non-finite
1462/// `rtol`, `Error::Extension` for unsupported integer or boolean dtypes,
1463/// numerical non-convergence, or a backend failure, and `Error::RuntimeState`
1464/// when the extension is not registered.
1465///
1466/// # Deferred errors
1467///
1468/// Symbolic shapes are materialized by this helper; failures are reported as
1469/// `ShapeConstraintEvaluation` or `ShapeExpressionEvaluation`.
1470pub fn pinv_with_rtol(a: &TracedTensor, rtol: f64) -> Result<TracedTensor> {
1471    ensure_float_or_complex("pinv_with_rtol", a.dtype)?;
1472    require_concrete_shape("pinv_with_rtol", a)?;
1473    let (u, s, vt) = svd(a)?;
1474    let abs_s = s.abs()?;
1475    let s_max = abs_s.reduce_max(Some(&[0]))?;
1476    let s_max_shape = s_max.concrete_shape()?;
1477    let threshold_scalar = broadcast_scalar(scalar_real(s.dtype, rtol.max(0.0))?, &s_max_shape)?;
1478    let threshold = (&s_max * &threshold_scalar)?;
1479    let s_shape = s.concrete_shape()?;
1480    let threshold = broadcast_batch_scalar_to_leading_axis(&threshold, &s_shape)?;
1481    let mask = abs_s.compare(&threshold, CompareDir::Gt)?;
1482    let mask = mask.convert(s.dtype)?;
1483    let ones = ones_like(&s)?;
1484    let neg_mask = (-&mask)?;
1485    let denom = (&s + &(&ones + &neg_mask)?)?;
1486    let s_inv = (&mask / &denom)?;
1487
1488    let v = vt.conj()?.transpose(&matrix_transpose_perm(vt.rank))?;
1489    let uh = u.conj()?.transpose(&matrix_transpose_perm(u.rank))?;
1490    let vs = scale_matrix_columns(&v, &s_inv)?;
1491    matmul_preserve_trailing_batch(&vs, &uh)
1492}
1493
1494/// Build a traced vector, matrix, or tensor norm op.
1495///
1496/// Floating-point and complex inputs are supported. Integer and boolean inputs
1497/// return an unsupported-dtype error.
1498///
1499/// # Examples
1500///
1501/// ```
1502/// use tenferro_linalg::TracedTensorLinalgExt;
1503/// use tenferro_runtime::TracedTensor;
1504///
1505/// let x = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap();
1506/// let length = x.norm(Some(2.0), Some(&[0]), false).unwrap();
1507/// assert_eq!(length.rank, 0);
1508/// ```
1509///
1510/// # Errors
1511///
1512/// Returns `Error::Validation` for an invalid axis, rank, or norm order,
1513/// `Error::Extension` for unsupported integer or boolean dtypes or a backend
1514/// numerical failure, and `Error::RuntimeState` when the extension is not
1515/// registered.
1516///
1517/// # Deferred errors
1518///
1519/// Backend numerical and runtime failures can occur during execution.
1520/// Symbolic input shapes are not deferred: this helper requires concrete
1521/// dimensions and returns `Error::TensorRuntime` wrapping
1522/// `ValidationError::InvalidArgument` for `shape` during graph construction,
1523/// regardless of `keepdim`.
1524pub fn norm(
1525    a: &TracedTensor,
1526    ord: Option<f64>,
1527    dim: Option<&[usize]>,
1528    keepdim: bool,
1529) -> Result<TracedTensor> {
1530    ensure_float_or_complex("norm", a.dtype)?;
1531    let shape = require_concrete_shape("norm", a)?;
1532    let axes = dim.map_or_else(|| (0..a.rank).collect::<Vec<_>>(), |dims| dims.to_vec());
1533    validate_axes("norm", a.rank, &axes)?;
1534    if axes.is_empty() {
1535        return Ok(a.clone());
1536    }
1537    if reduced_axes_have_zero_extent(&shape, &axes)
1538        && let Some(zero) = zero_norm_for_empty_reduction(a.dtype, &shape, &axes, keepdim, ord)?
1539    {
1540        return Ok(zero);
1541    }
1542
1543    let out = if can_square_without_abs(a.dtype, axes.len(), ord) {
1544        frobenius_norm(a, &axes)?
1545    } else {
1546        match axes.len() {
1547            1 => vector_norm(a, axes[0], ord)?,
1548            2 => matrix_norm(a, &axes, ord)?,
1549            _ => {
1550                let abs = a.abs()?;
1551                match ord {
1552                    None => frobenius_norm(&abs, &axes)?,
1553                    Some(p) if p == f64::INFINITY => abs.reduce_max(Some(&axes))?,
1554                    Some(p) if p == f64::NEG_INFINITY => abs.reduce_min(Some(&axes))?,
1555                    Some(0.0) => count_nonzero(&abs, &axes)?,
1556                    Some(p) => p_norm(&abs, &axes, p)?,
1557                }
1558            }
1559        }
1560    };
1561    restore_keepdim(out, &shape, &axes, keepdim)
1562}
1563
1564fn unexpected_output_count(name: &str, expected: usize) -> Error {
1565    Error::Internal(format!("{name} must produce exactly {expected} outputs"))
1566}
1567
1568fn one_output(outputs: Vec<TracedTensor>, name: &str) -> Result<TracedTensor> {
1569    let mut outputs = outputs.into_iter();
1570    match (outputs.next(), outputs.next()) {
1571        (Some(output), None) => Ok(output),
1572        _ => Err(unexpected_output_count(name, 1)),
1573    }
1574}
1575
1576fn two_outputs(outputs: Vec<TracedTensor>, name: &str) -> Result<(TracedTensor, TracedTensor)> {
1577    let mut outputs = outputs.into_iter();
1578    match (outputs.next(), outputs.next(), outputs.next()) {
1579        (Some(lhs), Some(rhs), None) => Ok((lhs, rhs)),
1580        _ => Err(unexpected_output_count(name, 2)),
1581    }
1582}
1583
1584fn three_outputs(
1585    outputs: Vec<TracedTensor>,
1586    name: &str,
1587) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
1588    let mut outputs = outputs.into_iter();
1589    match (
1590        outputs.next(),
1591        outputs.next(),
1592        outputs.next(),
1593        outputs.next(),
1594    ) {
1595        (Some(first), Some(second), Some(third), None) => Ok((first, second, third)),
1596        _ => Err(unexpected_output_count(name, 3)),
1597    }
1598}
1599
1600fn four_outputs(
1601    outputs: Vec<TracedTensor>,
1602    name: &str,
1603) -> Result<(TracedTensor, TracedTensor, TracedTensor, TracedTensor)> {
1604    let mut outputs = outputs.into_iter();
1605    match (
1606        outputs.next(),
1607        outputs.next(),
1608        outputs.next(),
1609        outputs.next(),
1610        outputs.next(),
1611    ) {
1612        (Some(first), Some(second), Some(third), Some(fourth), None) => {
1613            Ok((first, second, third, fourth))
1614        }
1615        _ => Err(unexpected_output_count(name, 4)),
1616    }
1617}
1618
1619fn five_outputs(
1620    outputs: Vec<TracedTensor>,
1621    name: &str,
1622) -> Result<(
1623    TracedTensor,
1624    TracedTensor,
1625    TracedTensor,
1626    TracedTensor,
1627    TracedTensor,
1628)> {
1629    let mut outputs = outputs.into_iter();
1630    match (
1631        outputs.next(),
1632        outputs.next(),
1633        outputs.next(),
1634        outputs.next(),
1635        outputs.next(),
1636        outputs.next(),
1637    ) {
1638        (Some(first), Some(second), Some(third), Some(fourth), Some(fifth), None) => {
1639            Ok((first, second, third, fourth, fifth))
1640        }
1641        _ => Err(unexpected_output_count(name, 5)),
1642    }
1643}
1644
1645fn scalar_real(dtype: DType, value: f64) -> Result<TracedTensor> {
1646    match dtype {
1647        DType::F64 => TracedTensor::from_vec_col_major(vec![], vec![value]),
1648        DType::F32 => TracedTensor::from_vec_col_major(vec![], vec![value as f32]),
1649        DType::I32 => TracedTensor::from_vec_col_major(vec![], vec![value.round() as i32]),
1650        DType::I64 => TracedTensor::from_vec_col_major(vec![], vec![value.round() as i64]),
1651        DType::Bool => TracedTensor::from_vec_col_major(vec![], vec![value != 0.0]),
1652        DType::C64 => TracedTensor::from_vec_col_major(vec![], vec![Complex64::new(value, 0.0)]),
1653        DType::C32 => {
1654            TracedTensor::from_vec_col_major(vec![], vec![Complex32::new(value as f32, 0.0)])
1655        }
1656        // An externally defined scalar has no traced scalar literal.
1657        DType::External(_) => Err(Error::TensorRuntime(
1658            tenferro_tensor::Error::invalid_argument(
1659                "scalar_real",
1660                "dtype",
1661                "an externally defined scalar has no traced scalar literal",
1662            ),
1663        )),
1664    }
1665}
1666
1667fn filled_real(dtype: DType, shape: Vec<usize>, value: f64) -> Result<TracedTensor> {
1668    let len = tenferro_tensor::validate::checked_shape_product("slogdet", "output shape", &shape)?;
1669    match dtype {
1670        DType::F64 => TracedTensor::from_vec_col_major(shape, vec![value; len]),
1671        DType::F32 => TracedTensor::from_vec_col_major(shape, vec![value as f32; len]),
1672        DType::I32 => TracedTensor::from_vec_col_major(shape, vec![value.round() as i32; len]),
1673        DType::I64 => TracedTensor::from_vec_col_major(shape, vec![value.round() as i64; len]),
1674        DType::Bool => TracedTensor::from_vec_col_major(shape, vec![value != 0.0; len]),
1675        DType::C64 => {
1676            TracedTensor::from_vec_col_major(shape, vec![Complex64::new(value, 0.0); len])
1677        }
1678        DType::C32 => {
1679            TracedTensor::from_vec_col_major(shape, vec![Complex32::new(value as f32, 0.0); len])
1680        }
1681        // An externally defined scalar has no traced filled literal.
1682        DType::External(_) => Err(Error::TensorRuntime(
1683            tenferro_tensor::Error::invalid_argument(
1684                "filled_real",
1685                "dtype",
1686                "an externally defined scalar has no traced filled literal",
1687            ),
1688        )),
1689    }
1690}
1691
1692fn real_values_dtype(dtype: DType) -> DType {
1693    match dtype {
1694        DType::C64 => DType::F64,
1695        DType::C32 => DType::F32,
1696        other => other,
1697    }
1698}
1699
1700fn can_square_without_abs(dtype: DType, axes_len: usize, ord: Option<f64>) -> bool {
1701    matches!(dtype, DType::F32 | DType::F64)
1702        && (ord.is_none() || (ord == Some(2.0) && axes_len != 2))
1703}
1704
1705fn ensure_min_rank(op: &'static str, actual: usize, expected: usize) -> Result<()> {
1706    if actual < expected {
1707        return Err(Error::TensorRuntime(tenferro_tensor::Error::rank_mismatch(
1708            op, expected, actual,
1709        )));
1710    }
1711    Ok(())
1712}
1713
1714fn validate_axes(op: &'static str, rank: usize, axes: &[usize]) -> Result<()> {
1715    tenferro_tensor::validate::validate_unique_axes(op, "dim", rank, axes)
1716        .map_err(Error::TensorRuntime)
1717}
1718
1719fn require_concrete_shape(op: &'static str, input: &TracedTensor) -> Result<Vec<usize>> {
1720    input.try_concrete_shape().ok_or_else(|| {
1721        Error::TensorRuntime(tenferro_tensor::Error::invalid_argument(
1722            op,
1723            "shape",
1724            "symbolic shape is not supported by this traced linalg helper",
1725        ))
1726    })
1727}
1728
1729fn zero_scalar(dtype: DType) -> Result<TracedTensor> {
1730    scalar_real(dtype, 0.0)
1731}
1732
1733fn one_scalar(dtype: DType) -> Result<TracedTensor> {
1734    scalar_real(dtype, 1.0)
1735}
1736
1737fn ones_like(input: &TracedTensor) -> Result<TracedTensor> {
1738    let shape = input.concrete_shape()?;
1739    broadcast_scalar(one_scalar(input.dtype)?, &shape)
1740}
1741
1742fn eye_like(anchor: &TracedTensor, size: usize) -> Result<TracedTensor> {
1743    let mut vector_shape = vec![size];
1744    let anchor_shape = anchor.concrete_shape()?;
1745    vector_shape.extend_from_slice(&anchor_shape[2..]);
1746    let diagonal = broadcast_scalar(one_scalar(anchor.dtype)?, &vector_shape)?;
1747    diagonal.embed_diag(0, 1)
1748}
1749
1750fn broadcast_scalar(input: TracedTensor, shape: &[usize]) -> Result<TracedTensor> {
1751    let input_shape = input.concrete_shape()?;
1752    if input_shape == shape {
1753        return Ok(input);
1754    }
1755    input.broadcast_in_dim(shape, &[])
1756}
1757
1758fn broadcast_batch_scalar_to_leading_axis(
1759    input: &TracedTensor,
1760    shape: &[usize],
1761) -> Result<TracedTensor> {
1762    let input_shape = input.concrete_shape()?;
1763    if input_shape == shape {
1764        return Ok(input.clone());
1765    }
1766    let dims: Vec<usize> = (1..shape.len()).collect();
1767    input.broadcast_in_dim(shape, &dims)
1768}
1769
1770fn matmul_preserve_trailing_batch(lhs: &TracedTensor, rhs: &TracedTensor) -> Result<TracedTensor> {
1771    let rank = lhs.rank;
1772    let batch_dims: Vec<usize> = (2..rank).collect();
1773    lhs.dot_general(
1774        rhs,
1775        DotGeneralConfig {
1776            lhs_contracting_dims: [1].as_slice().into(),
1777            rhs_contracting_dims: [0].as_slice().into(),
1778            lhs_batch_dims: batch_dims.clone().into(),
1779            rhs_batch_dims: batch_dims.into(),
1780        },
1781    )
1782}
1783
1784fn matrix_transpose_perm(rank: usize) -> Vec<usize> {
1785    let mut perm: Vec<usize> = (0..rank).collect();
1786    perm.swap(0, 1);
1787    perm
1788}
1789
1790fn frobenius_norm(abs: &TracedTensor, axes: &[usize]) -> Result<TracedTensor> {
1791    abs.reduce_sum_squares(Some(axes))?.sqrt()
1792}
1793
1794fn p_norm(abs: &TracedTensor, axes: &[usize], p: f64) -> Result<TracedTensor> {
1795    if !p.is_finite() || p == 0.0 {
1796        return Err(Error::invalid_argument(
1797            "norm",
1798            ErrorPhase::GraphBuild,
1799            "p",
1800            format!("p-norm order must be finite and nonzero, got {p}"),
1801        ));
1802    }
1803    if p == 2.0 {
1804        return frobenius_norm(abs, axes);
1805    }
1806    let power = abs.pow(&scalar_real(abs.dtype, p)?)?;
1807    let inv_p = scalar_real(abs.dtype, 1.0 / p)?;
1808    power.reduce_sum(Some(axes))?.pow(&inv_p)
1809}
1810
1811fn reduced_axes_have_zero_extent(shape: &[usize], axes: &[usize]) -> bool {
1812    axes.iter().any(|&axis| shape[axis] == 0)
1813}
1814
1815fn zero_norm_for_empty_reduction(
1816    dtype: DType,
1817    input_shape: &[usize],
1818    axes: &[usize],
1819    keepdim: bool,
1820    ord: Option<f64>,
1821) -> Result<Option<TracedTensor>> {
1822    if !empty_reduction_norm_is_zero(axes.len(), ord) {
1823        return Ok(None);
1824    }
1825    let output_shape = reduction_shape(input_shape, axes, keepdim);
1826    zero_traced_tensor(real_norm_dtype(dtype)?, output_shape).map(Some)
1827}
1828
1829fn empty_reduction_norm_is_zero(axis_count: usize, ord: Option<f64>) -> bool {
1830    match ord {
1831        None => true,
1832        Some(0.0) => true,
1833        Some(p) if p.is_infinite() => true,
1834        Some(p) if p.is_finite() && p > 0.0 => axis_count != 2 || p != 2.0,
1835        _ => false,
1836    }
1837}
1838
1839fn reduction_shape(input_shape: &[usize], axes: &[usize], keepdim: bool) -> Vec<usize> {
1840    if keepdim {
1841        let mut shape = input_shape.to_vec();
1842        for &axis in axes {
1843            shape[axis] = 1;
1844        }
1845        return shape;
1846    }
1847    let mut reduced = vec![false; input_shape.len()];
1848    for &axis in axes {
1849        reduced[axis] = true;
1850    }
1851    input_shape
1852        .iter()
1853        .enumerate()
1854        .filter_map(|(axis, &dim)| (!reduced[axis]).then_some(dim))
1855        .collect()
1856}
1857
1858fn real_norm_dtype(dtype: DType) -> Result<DType> {
1859    match dtype {
1860        DType::F32 | DType::F64 => Ok(dtype),
1861        DType::C32 => Ok(DType::F32),
1862        DType::C64 => Ok(DType::F64),
1863        _ => Err(Error::TensorRuntime(
1864            tenferro_tensor::Error::unsupported_dtype(
1865                "norm",
1866                dtype,
1867                "norm supports only floating-point and complex dtypes",
1868            ),
1869        )),
1870    }
1871}
1872
1873fn zero_traced_tensor(dtype: DType, shape: Vec<usize>) -> Result<TracedTensor> {
1874    let len = checked_element_count("norm", &shape)?;
1875    match dtype {
1876        DType::F32 => TracedTensor::from_vec_col_major(shape, vec![0.0_f32; len]),
1877        DType::F64 => TracedTensor::from_vec_col_major(shape, vec![0.0_f64; len]),
1878        DType::C32 => TracedTensor::from_vec_col_major(shape, vec![Complex32::new(0.0, 0.0); len]),
1879        DType::C64 => TracedTensor::from_vec_col_major(shape, vec![Complex64::new(0.0, 0.0); len]),
1880        _ => Err(Error::TensorRuntime(
1881            tenferro_tensor::Error::unsupported_dtype(
1882                "norm",
1883                dtype,
1884                "norm supports only floating-point and complex dtypes",
1885            ),
1886        )),
1887    }
1888}
1889
1890fn checked_element_count(op: &'static str, shape: &[usize]) -> Result<usize> {
1891    shape.iter().try_fold(1usize, |acc, &dim| {
1892        acc.checked_mul(dim).ok_or_else(|| {
1893            Error::TensorRuntime(tenferro_tensor::Error::invalid_argument(
1894                op,
1895                "shape",
1896                "shape element count overflow",
1897            ))
1898        })
1899    })
1900}
1901
1902fn default_pinv_rtol(dtype: DType, max_dim: usize) -> f64 {
1903    let eps = match dtype {
1904        DType::F32 | DType::C32 => f32::EPSILON as f64,
1905        DType::F64 | DType::C64 => f64::EPSILON,
1906        DType::I32 | DType::I64 | DType::Bool => 0.0,
1907        // INVARIANT: the decomposition rejects an externally defined scalar before
1908        // it asks for a tolerance, so this value is never used.
1909        DType::External(_) => unreachable!("the decomposition validates its dtype first"),
1910    };
1911    eps * max_dim as f64
1912}
1913
1914fn vector_norm(a: &TracedTensor, axis: usize, ord: Option<f64>) -> Result<TracedTensor> {
1915    let abs = a.abs()?;
1916    match ord {
1917        None => frobenius_norm(&abs, &[axis]),
1918        Some(0.0) => count_nonzero(&abs, &[axis]),
1919        Some(p) if p == f64::INFINITY => abs.reduce_max(Some(&[axis])),
1920        Some(p) if p == f64::NEG_INFINITY => abs.reduce_min(Some(&[axis])),
1921        Some(p) => p_norm(&abs, &[axis], p),
1922    }
1923}
1924
1925fn matrix_norm(a: &TracedTensor, axes: &[usize], ord: Option<f64>) -> Result<TracedTensor> {
1926    let matrix = move_axes_to_front(a, axes)?;
1927    if matches!(ord, Some(2.0) | Some(-2.0)) {
1928        let singular_values = svd_values(&matrix)?.abs()?;
1929        return if ord == Some(2.0) {
1930            singular_values.reduce_max(Some(&[0]))
1931        } else {
1932            singular_values.reduce_min(Some(&[0]))
1933        };
1934    }
1935
1936    let abs = matrix.abs()?;
1937    match ord {
1938        None => frobenius_norm(&abs, &[0, 1]),
1939        Some(p) if p == f64::INFINITY => matrix_row_sum_norm(&abs, true),
1940        Some(p) if p == f64::NEG_INFINITY => matrix_row_sum_norm(&abs, false),
1941        Some(1.0) => matrix_col_sum_norm(&abs, true),
1942        Some(-1.0) => matrix_col_sum_norm(&abs, false),
1943        Some(0.0) => count_nonzero(&abs, &[0, 1]),
1944        Some(p) => p_norm(&abs, &[0, 1], p),
1945    }
1946}
1947
1948fn svd_values(a: &TracedTensor) -> Result<TracedTensor> {
1949    let (_u, s, _vt) = three_outputs(
1950        apply(
1951            Arc::new(LinalgExtensionOp::new(LinalgOp::Svd {
1952                derivative_eps: SvdOptions::default().derivative_eps,
1953                gauge: SvdOptions::default().gauge,
1954                driver: SvdOptions::default().driver,
1955            })),
1956            &[a],
1957        )?,
1958        "svd_values",
1959    )?;
1960    Ok(s)
1961}
1962
1963fn eigh_values(a: &TracedTensor) -> Result<TracedTensor> {
1964    let (values, _vectors) = two_outputs(
1965        apply(
1966            Arc::new(LinalgExtensionOp::new(LinalgOp::Eigh {
1967                derivative_eps: EighOptions::default().derivative_eps,
1968                gauge: EighOptions::default().gauge,
1969                driver: EighOptions::default().driver,
1970            })),
1971            &[a],
1972        )?,
1973        "eigh_values",
1974    )?;
1975    Ok(values)
1976}
1977
1978fn eig_values(a: &TracedTensor) -> Result<TracedTensor> {
1979    let (values, _vectors) = two_outputs(
1980        apply(
1981            Arc::new(LinalgExtensionOp::new(LinalgOp::Eig {
1982                input_dtype: a.dtype,
1983            })),
1984            &[a],
1985        )?,
1986        "eig_values",
1987    )?;
1988    Ok(values)
1989}
1990
1991fn scale_matrix_columns(matrix: &TracedTensor, scale: &TracedTensor) -> Result<TracedTensor> {
1992    let matrix_shape = matrix.concrete_shape()?;
1993    let scale_shape_input = scale.concrete_shape()?;
1994    let mut scale_shape = vec![1, scale_shape_input[0]];
1995    scale_shape.extend_from_slice(&matrix_shape[2..]);
1996    let dims: Vec<usize> = (0..matrix_shape.len()).collect();
1997    let scale = scale
1998        .reshape(&scale_shape)?
1999        .broadcast_in_dim(&matrix_shape, &dims)?;
2000    matrix * &scale
2001}
2002
2003fn count_nonzero(abs: &TracedTensor, axes: &[usize]) -> Result<TracedTensor> {
2004    let mask = abs.compare(&zero_scalar(abs.dtype)?, CompareDir::Gt)?;
2005    mask.convert(abs.dtype)?.reduce_sum(Some(axes))
2006}
2007
2008fn matrix_row_sum_norm(abs: &TracedTensor, take_max: bool) -> Result<TracedTensor> {
2009    let row_sums = abs.reduce_sum(Some(&[1]))?;
2010    if take_max {
2011        row_sums.reduce_max(Some(&[0]))
2012    } else {
2013        row_sums.reduce_min(Some(&[0]))
2014    }
2015}
2016
2017fn matrix_col_sum_norm(abs: &TracedTensor, take_max: bool) -> Result<TracedTensor> {
2018    let col_sums = abs.reduce_sum(Some(&[0]))?;
2019    if take_max {
2020        col_sums.reduce_max(Some(&[0]))
2021    } else {
2022        col_sums.reduce_min(Some(&[0]))
2023    }
2024}
2025
2026fn move_axes_to_front(tensor: &TracedTensor, axes: &[usize]) -> Result<TracedTensor> {
2027    if axes.iter().enumerate().all(|(index, &axis)| index == axis) {
2028        return Ok(tensor.clone());
2029    }
2030
2031    let mut selected = vec![false; tensor.rank];
2032    for &axis in axes {
2033        selected[axis] = true;
2034    }
2035
2036    let mut perm = Vec::with_capacity(tensor.rank);
2037    perm.extend_from_slice(axes);
2038    for (axis, is_selected) in selected.iter().enumerate().take(tensor.rank) {
2039        if !*is_selected {
2040            perm.push(axis);
2041        }
2042    }
2043    tensor.transpose(&perm)
2044}
2045
2046fn restore_keepdim(
2047    reduced: TracedTensor,
2048    original_shape: &[usize],
2049    axes: &[usize],
2050    keepdim: bool,
2051) -> Result<TracedTensor> {
2052    if !keepdim {
2053        return Ok(reduced);
2054    }
2055    let mut kept_shape = original_shape.to_vec();
2056    for &axis in axes {
2057        kept_shape[axis] = 1;
2058    }
2059    reduced.reshape(&kept_shape)
2060}
2061
2062#[cfg(test)]
2063mod tests {
2064    use super::p_norm;
2065    use tenferro_runtime::TracedTensor;
2066
2067    #[test]
2068    fn p_norm_rejects_zero_and_non_finite_orders() {
2069        let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
2070        let abs = x.abs().unwrap();
2071
2072        for p in [0.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
2073            let err = p_norm(&abs, &[0], p).unwrap_err();
2074            assert!(
2075                err.to_string().contains("finite") || err.to_string().contains("nonzero"),
2076                "expected finite nonzero order error, got {err:?}"
2077            );
2078        }
2079    }
2080}