Skip to main content

tenferro_linalg/
extension.rs

1use std::any::Any;
2use std::hash::{Hash, Hasher};
3use std::sync::Arc;
4
5use num_complex::{Complex32, Complex64};
6use tenferro_cpu::with_cpu_exec_session;
7use tenferro_extension_macros::define_extension_runtime;
8use tenferro_ops::SymDim;
9use tenferro_runtime::extension::ExtensionOp;
10use tenferro_tensor::{BackendSession, DType, Error, ErrorKind, Tensor, TensorBackend, TensorRead};
11
12#[cfg(feature = "cuda")]
13use tenferro_gpu::cuda::with_cuda_exec_session;
14
15use crate::backend::LinalgBackend;
16use crate::RankRevealingQrOptions;
17
18#[cfg(all(test, feature = "cuda"))]
19#[path = "extension/cuda_tests.rs"]
20mod cuda_tests;
21mod gauge;
22#[cfg(all(test, not(feature = "cuda")))]
23mod tests;
24
25pub(crate) use gauge::{apply_eigh_gauge, apply_qr_gauge};
26
27pub const LINALG_EXTENSION_FAMILY_ID: &str = "tenferro-linalg.linalg.v1";
28
29/// Default derivative regularization used by decomposition AD rules.
30///
31/// This epsilon is used only when differentiating decomposition formulas with
32/// repeated or nearly repeated spectral values. It is not a solver tolerance.
33///
34/// # Examples
35///
36/// ```rust
37/// use tenferro_linalg::{SvdOptions, DEFAULT_DECOMPOSITION_DERIVATIVE_EPS};
38///
39/// let options = SvdOptions::default();
40/// assert_eq!(options.derivative_eps, DEFAULT_DECOMPOSITION_DERIVATIVE_EPS);
41/// ```
42pub const DEFAULT_DECOMPOSITION_DERIVATIVE_EPS: f64 = 1e-12;
43
44/// Singular-vector gauge convention used by [`SvdOptions`].
45///
46/// # Examples
47///
48/// ```rust
49/// use tenferro_linalg::{SvdGauge, SvdOptions};
50///
51/// let options = SvdOptions::default().gauge(SvdGauge::CanonicalPivot);
52/// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
53/// ```
54#[derive(Clone, Copy, Debug, PartialEq, Eq)]
55pub enum SvdGauge {
56    /// Leave the backend's raw singular vector signs or phases unchanged.
57    Raw,
58    /// Make each left singular vector's max-absolute pivot entry positive-real
59    /// and adjust the matching `VT` row so reconstruction is preserved.
60    CanonicalPivot,
61}
62
63/// Eigenvector gauge convention used by [`EighOptions`].
64///
65/// # Examples
66///
67/// ```rust
68/// use tenferro_linalg::{EighGauge, EighOptions};
69///
70/// let options = EighOptions::default().gauge(EighGauge::CanonicalPivot);
71/// assert_eq!(options.gauge, EighGauge::CanonicalPivot);
72/// ```
73#[derive(Clone, Copy, Debug, PartialEq, Eq)]
74pub enum EighGauge {
75    /// Leave the backend's raw eigenvector signs or phases unchanged.
76    Raw,
77    /// Make each eigenvector's max-absolute pivot entry positive-real.
78    CanonicalPivot,
79}
80
81/// QR factor gauge convention used by [`QrOptions`].
82///
83/// # Examples
84///
85/// ```rust
86/// use tenferro_linalg::{QrGauge, QrOptions};
87///
88/// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
89/// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
90/// ```
91#[derive(Clone, Copy, Debug, PartialEq, Eq)]
92pub enum QrGauge {
93    /// Leave the backend's raw QR signs or phases unchanged.
94    Raw,
95    /// Make each `R` diagonal entry positive-real, compensating `Q`.
96    PositiveDiagonal,
97}
98
99/// cuSOLVER SVD routine selection used by [`SvdOptions`] on the CUDA backend.
100///
101/// The driver changes speed and numerical behavior; [`Self::Xgesvdp`] may
102/// perturb near-singular inputs. Gauge handling and AD rule selection are
103/// unchanged. CPU providers have a single SVD kernel and ignore the driver.
104/// Equivalent to `jax.lax.linalg.svd(..., algorithm=...)` and
105/// `torch.linalg.svd(..., driver=...)`.
106///
107/// # Examples
108///
109/// ```rust
110/// use tenferro_linalg::{SvdDriver, SvdOptions};
111///
112/// let options = SvdOptions::default().driver(SvdDriver::Gesvd);
113/// assert_eq!(options.driver, SvdDriver::Gesvd);
114/// assert_eq!(SvdOptions::default().driver, SvdDriver::Auto);
115/// ```
116#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
117pub enum SvdDriver {
118    /// The backend's default policy. On CUDA this is the JAX-compatible rule:
119    /// `gesvdj` when both matrix dimensions are at most 1024, otherwise
120    /// `gesvd`.
121    #[default]
122    Auto,
123    /// cuSOLVER's Jacobi driver (`cusolverDn<t>gesvdj`) regardless of size.
124    Gesvdj,
125    /// cuSOLVER's QR-based driver (`cusolverDn<t>gesvd`) regardless of size.
126    Gesvd,
127    /// cuSOLVER's polar-decomposition driver (`cusolverDnXgesvdp`).
128    ///
129    /// Explicit only: [`Self::Auto`] never selects this routine. Near-singular
130    /// inputs may be perturbed by cuSOLVER, shifting small singular values.
131    /// The perturbation magnitude (`h_err_sigma`) is discarded, not returned
132    /// with the factors. CPU providers accept and ignore this selection.
133    ///
134    /// Accuracy is input-dependent, not uniformly better or worse than
135    /// [`Self::Gesvd`]. On ten-decade spectra built as `U diag(s) Vá´´` it was
136    /// the most accurate of the three drivers on an A100, but on a
137    /// circulant-like matrix with the same spectrum it lost about two digits
138    /// against `gesvd`. Measure on your own inputs before relying on it for
139    /// small singular values.
140    Xgesvdp,
141}
142
143/// Options for singular value decomposition.
144///
145/// # Examples
146///
147/// ```rust
148/// use tenferro_linalg::{SvdDriver, SvdGauge, SvdOptions};
149///
150/// let options = SvdOptions::default()
151///     .gauge(SvdGauge::CanonicalPivot)
152///     .derivative_eps(1.0e-10)
153///     .driver(SvdDriver::Gesvd);
154/// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
155/// assert_eq!(options.derivative_eps, 1.0e-10);
156/// assert_eq!(options.driver, SvdDriver::Gesvd);
157/// ```
158#[derive(Clone, Copy, Debug, PartialEq)]
159pub struct SvdOptions {
160    /// Singular-vector gauge convention.
161    pub gauge: SvdGauge,
162    /// AD derivative regularization for repeated or nearly repeated singular values.
163    pub derivative_eps: f64,
164    /// CUDA SVD driver; ignored by CPU providers.
165    pub driver: SvdDriver,
166}
167
168impl Default for SvdOptions {
169    fn default() -> Self {
170        Self {
171            gauge: SvdGauge::Raw,
172            derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
173            driver: SvdDriver::Auto,
174        }
175    }
176}
177
178impl SvdOptions {
179    /// Return options with the requested singular-vector gauge.
180    ///
181    /// # Examples
182    ///
183    /// ```rust
184    /// use tenferro_linalg::{SvdGauge, SvdOptions};
185    ///
186    /// let options = SvdOptions::default().gauge(SvdGauge::CanonicalPivot);
187    /// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
188    /// ```
189    pub fn gauge(mut self, gauge: SvdGauge) -> Self {
190        self.gauge = gauge;
191        self
192    }
193
194    /// Return options with an explicit derivative epsilon.
195    ///
196    /// # Examples
197    ///
198    /// ```rust
199    /// use tenferro_linalg::SvdOptions;
200    ///
201    /// let options = SvdOptions::default().derivative_eps(1.0e-9);
202    /// assert_eq!(options.derivative_eps, 1.0e-9);
203    /// ```
204    pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
205        self.derivative_eps = derivative_eps;
206        self
207    }
208
209    /// Return options with an explicit CUDA SVD driver.
210    ///
211    /// # Examples
212    ///
213    /// ```rust
214    /// use tenferro_linalg::{SvdDriver, SvdOptions};
215    ///
216    /// let options = SvdOptions::default().driver(SvdDriver::Gesvdj);
217    /// assert_eq!(options.driver, SvdDriver::Gesvdj);
218    /// ```
219    pub fn driver(mut self, driver: SvdDriver) -> Self {
220        self.driver = driver;
221        self
222    }
223}
224
225/// cuSOLVER Hermitian eigensolver selection used by [`EighOptions`] on the
226/// CUDA backend.
227///
228/// The driver changes speed and rounding, not the decomposition contract, so
229/// gauges and AD rules are unaffected. CPU providers have a single eigh kernel
230/// and ignore it. Analogous to [`SvdDriver`] and to
231/// `scipy.linalg.eigh(..., driver=...)`.
232///
233/// # Examples
234///
235/// ```rust
236/// use tenferro_linalg::{EighDriver, EighOptions};
237///
238/// let options = EighOptions::default().driver(EighDriver::Syevj);
239/// assert_eq!(options.driver, EighDriver::Syevj);
240/// assert_eq!(EighOptions::default().driver, EighDriver::Auto);
241/// ```
242#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
243pub enum EighDriver {
244    /// The backend's default policy: the fastest divide-and-conquer routine
245    /// available. On CUDA a single matrix uses `syevd`/`heevd` and a batch
246    /// uses `cusolverDnXsyevBatched`, which solves the whole batch in one
247    /// launch at the same accuracy. On an A100 that is about 100x to 280x
248    /// faster than the per-matrix loop this default used before.
249    ///
250    /// The batched routine is a different cuSOLVER entry point, so batched
251    /// results can differ from the per-matrix loop in the last ULPs.
252    #[default]
253    Auto,
254    /// cuSOLVER's divide-and-conquer driver, which is what [`Self::Auto`]
255    /// already selects. Kept as an explicit spelling of the same choice.
256    Syevd,
257    /// cuSOLVER's Jacobi driver (`cusolverDn<t>syevj`, `heevj` for complex
258    /// input, and their batched forms).
259    ///
260    /// Worth asking for only on batches of small matrices: measured on an
261    /// A100 it is slightly ahead of the default up to about order 32 and
262    /// loses from there, by about 4x at order 128 and far more on wide
263    /// spectra, where it is also the less accurate of the two.
264    Syevj,
265}
266
267/// Options for Hermitian eigenvalue decomposition.
268///
269/// # Examples
270///
271/// ```rust
272/// use tenferro_linalg::{EighDriver, EighOptions};
273///
274/// let options = EighOptions::default()
275///     .derivative_eps(1.0e-10)
276///     .driver(EighDriver::Syevj);
277/// assert_eq!(options.derivative_eps, 1.0e-10);
278/// assert_eq!(options.driver, EighDriver::Syevj);
279/// ```
280#[derive(Clone, Copy, Debug, PartialEq)]
281pub struct EighOptions {
282    /// Eigenvector gauge convention.
283    pub gauge: EighGauge,
284    /// AD derivative regularization for repeated or nearly repeated eigenvalues.
285    pub derivative_eps: f64,
286    /// CUDA eigensolver driver; ignored by CPU providers.
287    pub driver: EighDriver,
288}
289
290impl Default for EighOptions {
291    fn default() -> Self {
292        Self {
293            gauge: EighGauge::Raw,
294            derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
295            driver: EighDriver::Auto,
296        }
297    }
298}
299
300impl EighOptions {
301    /// Return options with the requested eigenvector gauge.
302    ///
303    /// # Examples
304    ///
305    /// ```rust
306    /// use tenferro_linalg::{EighGauge, EighOptions};
307    ///
308    /// let options = EighOptions::default().gauge(EighGauge::CanonicalPivot);
309    /// assert_eq!(options.gauge, EighGauge::CanonicalPivot);
310    /// ```
311    pub fn gauge(mut self, gauge: EighGauge) -> Self {
312        self.gauge = gauge;
313        self
314    }
315
316    /// Return options with an explicit derivative epsilon.
317    ///
318    /// # Examples
319    ///
320    /// ```rust
321    /// use tenferro_linalg::EighOptions;
322    ///
323    /// let options = EighOptions::default().derivative_eps(1.0e-9);
324    /// assert_eq!(options.derivative_eps, 1.0e-9);
325    /// ```
326    pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
327        self.derivative_eps = derivative_eps;
328        self
329    }
330
331    /// Return options with an explicit CUDA eigensolver driver.
332    ///
333    /// # Examples
334    ///
335    /// ```rust
336    /// use tenferro_linalg::{EighDriver, EighOptions};
337    ///
338    /// let options = EighOptions::default().driver(EighDriver::Syevd);
339    /// assert_eq!(options.driver, EighDriver::Syevd);
340    /// ```
341    pub fn driver(mut self, driver: EighDriver) -> Self {
342        self.driver = driver;
343        self
344    }
345}
346
347/// Options for QR decomposition.
348///
349/// # Examples
350///
351/// ```rust
352/// use tenferro_linalg::{QrGauge, QrOptions};
353///
354/// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
355/// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
356/// ```
357#[derive(Clone, Copy, Debug, PartialEq, Eq)]
358pub struct QrOptions {
359    /// QR sign or phase convention.
360    pub gauge: QrGauge,
361}
362
363impl Default for QrOptions {
364    fn default() -> Self {
365        Self {
366            gauge: QrGauge::Raw,
367        }
368    }
369}
370
371impl QrOptions {
372    /// Return options with the requested QR gauge.
373    ///
374    /// # Examples
375    ///
376    /// ```rust
377    /// use tenferro_linalg::{QrGauge, QrOptions};
378    ///
379    /// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
380    /// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
381    /// ```
382    pub fn gauge(mut self, gauge: QrGauge) -> Self {
383        self.gauge = gauge;
384        self
385    }
386}
387
388pub(crate) fn validate_derivative_eps(
389    op: &'static str,
390    derivative_eps: f64,
391) -> tenferro_tensor::Result<()> {
392    if derivative_eps.is_finite() && derivative_eps > 0.0 {
393        Ok(())
394    } else {
395        Err(Error::invalid_argument(
396            op,
397            "derivative_eps",
398            format!("must be positive and finite, got {derivative_eps}"),
399        ))
400    }
401}
402
403#[derive(Clone, Copy, Debug, PartialEq)]
404#[doc(hidden)]
405#[allow(dead_code)]
406pub(crate) enum LinalgOp {
407    Cholesky,
408    Lu,
409    LuFactor,
410    LuSolvePrepared {
411        transpose_a: bool,
412        conjugate_a: bool,
413    },
414    SignDetFromLuFactor,
415    LogAbsDetFromLuFactor,
416    FullPivLu,
417    FullPivLuSolve {
418        transpose_a: bool,
419    },
420    /// Solve `a @ x = b` with partial-pivot LU (same kernel as
421    /// `LinalgBackend::solve`). Two inputs (matrix, rhs) to one output.
422    /// The untracked eager surface constructs this variant.
423    Solve,
424    /// Fused partial-pivot solve that also returns its factors: inputs
425    /// `(a, b)`, outputs `(x, packed_lu, pivots)`, with the `LuFactor` packed
426    /// layout and 1-based LAPACK pivots.
427    ///
428    /// Tracked eager and traced `solve` emit this op so one backend call both
429    /// factors and solves while the factors stay available as AD residuals:
430    /// the tangent and the transpose solve reuse them through
431    /// [`LinalgOp::LuSolvePrepared`] instead of refactoring. It is never
432    /// pruned to [`LinalgOp::Solve`]; see `prune_outputs`.
433    LuFactorSolve,
434    Svd {
435        derivative_eps: f64,
436        gauge: SvdGauge,
437        driver: SvdDriver,
438    },
439    /// Full-matrices SVD: `U` is `m x m` and `Vh` is `n x n`, so the trailing
440    /// `Vh` rows span the input's right nullspace. Value-only: AD is
441    /// intentionally unsupported (see the linalg AD support manifest).
442    SvdFull,
443    /// Singular values only.
444    SvdVals {
445        derivative_eps: f64,
446        driver: SvdDriver,
447    },
448    Qr {
449        gauge: QrGauge,
450    },
451    RankRevealingQr {
452        gauge: QrGauge,
453        rtol: f64,
454        atol: f64,
455    },
456    HouseholderQrFactor,
457    HouseholderQrFromFactors,
458    HouseholderQrAppend,
459    HouseholderQrR {
460        gauge: QrGauge,
461    },
462    HouseholderQrQColumns {
463        start: usize,
464        end: usize,
465        gauge: QrGauge,
466    },
467    /// Internal AD residual operation for symbolic full thin-Q recovery.
468    HouseholderQrThinQ {
469        gauge: QrGauge,
470    },
471    /// Internal linear operation for abstract-state column append.
472    HouseholderQrAppendTangent,
473    /// Internal transpose operation that splits an appended state cotangent.
474    HouseholderQrSplitTangent {
475        right: bool,
476    },
477    Eigh {
478        derivative_eps: f64,
479        gauge: EighGauge,
480        driver: EighDriver,
481    },
482    /// Eigenvalues only.
483    EighVals {
484        derivative_eps: f64,
485        driver: EighDriver,
486    },
487    Eig {
488        input_dtype: DType,
489    },
490    EigVals {
491        input_dtype: DType,
492    },
493    TriangularSolve {
494        left_side: bool,
495        lower: bool,
496        transpose_a: bool,
497        unit_diagonal: bool,
498    },
499}
500
501impl LinalgOp {
502    fn output_count(self) -> usize {
503        match self {
504            Self::Cholesky
505            | Self::EighVals { .. }
506            | Self::EigVals { .. }
507            | Self::FullPivLuSolve { .. }
508            | Self::LogAbsDetFromLuFactor
509            | Self::LuSolvePrepared { .. }
510            | Self::SignDetFromLuFactor
511            | Self::Solve
512            | Self::SvdVals { .. }
513            | Self::TriangularSolve { .. } => 1,
514            Self::Svd { .. } | Self::SvdFull | Self::LuFactorSolve => 3,
515            Self::RankRevealingQr { .. } | Self::Lu => 4,
516            Self::Qr { .. }
517            | Self::HouseholderQrFactor
518            | Self::HouseholderQrFromFactors
519            | Self::HouseholderQrAppend
520            | Self::Eigh { .. }
521            | Self::Eig { .. } => 2,
522            Self::HouseholderQrR { .. }
523            | Self::HouseholderQrQColumns { .. }
524            | Self::HouseholderQrThinQ { .. }
525            | Self::HouseholderQrAppendTangent
526            | Self::HouseholderQrSplitTangent { .. } => 1,
527            Self::LuFactor => 3,
528            Self::FullPivLu => 5,
529        }
530    }
531
532    fn input_count(self) -> usize {
533        match self {
534            Self::FullPivLuSolve { .. }
535            | Self::Solve
536            | Self::LuFactorSolve
537            | Self::TriangularSolve { .. }
538            | Self::HouseholderQrFromFactors
539            | Self::HouseholderQrR { .. }
540            | Self::HouseholderQrQColumns { .. }
541            | Self::HouseholderQrThinQ { .. } => 2,
542            Self::LogAbsDetFromLuFactor => 2,
543            Self::SignDetFromLuFactor | Self::HouseholderQrAppend => 3,
544            Self::HouseholderQrAppendTangent => 4,
545            Self::HouseholderQrSplitTangent { .. } => 3,
546            Self::LuSolvePrepared { .. } => 4,
547            _ => 1,
548        }
549    }
550
551    fn tag(self) -> u8 {
552        match self {
553            Self::Cholesky => 0,
554            Self::Lu => 1,
555            Self::FullPivLu => 2,
556            Self::FullPivLuSolve { .. } => 3,
557            Self::Svd { .. } => 4,
558            Self::Qr { .. } => 5,
559            Self::Eigh { .. } => 6,
560            Self::Eig { .. } => 7,
561            Self::TriangularSolve { .. } => 9,
562            Self::LuFactor => 10,
563            Self::LuSolvePrepared { .. } => 11,
564            Self::SvdVals { .. } => 12,
565            Self::EighVals { .. } => 13,
566            Self::EigVals { .. } => 14,
567            Self::SvdFull => 15,
568            Self::LogAbsDetFromLuFactor => 16,
569            Self::SignDetFromLuFactor => 17,
570            Self::Solve => 18,
571            Self::HouseholderQrFactor => 19,
572            Self::HouseholderQrFromFactors => 20,
573            Self::HouseholderQrAppend => 21,
574            Self::HouseholderQrR { .. } => 22,
575            Self::HouseholderQrQColumns { .. } => 23,
576            Self::HouseholderQrThinQ { .. } => 24,
577            Self::HouseholderQrAppendTangent => 25,
578            Self::HouseholderQrSplitTangent { .. } => 26,
579            Self::RankRevealingQr { .. } => 27,
580            Self::LuFactorSolve => 28,
581        }
582    }
583}
584
585#[derive(Clone, Debug, PartialEq)]
586#[doc(hidden)]
587pub(crate) struct LinalgExtensionOp {
588    op: LinalgOp,
589}
590
591impl LinalgExtensionOp {
592    pub(crate) fn new(op: LinalgOp) -> Self {
593        Self { op }
594    }
595
596    pub(crate) fn op(&self) -> LinalgOp {
597        self.op
598    }
599}
600
601impl ExtensionOp for LinalgExtensionOp {
602    fn family_id(&self) -> &'static str {
603        LINALG_EXTENSION_FAMILY_ID
604    }
605
606    fn payload_hash(&self, hasher: &mut dyn Hasher) {
607        hasher.write_u8(self.op.tag());
608        match self.op {
609            LinalgOp::Svd {
610                derivative_eps,
611                gauge,
612                driver,
613            } => {
614                hasher.write_u64(derivative_eps.to_bits());
615                hash_svd_gauge(hasher, gauge);
616                hash_svd_driver(hasher, driver);
617            }
618            LinalgOp::SvdVals {
619                derivative_eps,
620                driver,
621            } => {
622                hasher.write_u64(derivative_eps.to_bits());
623                hash_svd_driver(hasher, driver);
624            }
625            LinalgOp::EighVals {
626                derivative_eps,
627                driver,
628            } => {
629                hasher.write_u64(derivative_eps.to_bits());
630                hash_eigh_driver(hasher, driver);
631            }
632            LinalgOp::Qr { gauge }
633            | LinalgOp::HouseholderQrR { gauge }
634            | LinalgOp::HouseholderQrThinQ { gauge } => {
635                hash_qr_gauge(hasher, gauge);
636            }
637            LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
638                hash_qr_gauge(hasher, gauge);
639                hasher.write_u64(rtol.to_bits());
640                hasher.write_u64(atol.to_bits());
641            }
642            LinalgOp::HouseholderQrQColumns { start, end, gauge } => {
643                hasher.write_usize(start);
644                hasher.write_usize(end);
645                hash_qr_gauge(hasher, gauge);
646            }
647            LinalgOp::Eigh {
648                derivative_eps,
649                gauge,
650                driver,
651            } => {
652                hasher.write_u64(derivative_eps.to_bits());
653                hash_eigh_gauge(hasher, gauge);
654                hash_eigh_driver(hasher, driver);
655            }
656            LinalgOp::Eig { input_dtype } | LinalgOp::EigVals { input_dtype } => {
657                hash_dtype(hasher, input_dtype);
658            }
659            LinalgOp::FullPivLuSolve { transpose_a }
660            | LinalgOp::HouseholderQrSplitTangent { right: transpose_a } => {
661                hasher.write_u8(u8::from(transpose_a));
662            }
663            LinalgOp::LuSolvePrepared {
664                transpose_a,
665                conjugate_a,
666            } => {
667                hasher.write_u8(u8::from(transpose_a));
668                hasher.write_u8(u8::from(conjugate_a));
669            }
670            LinalgOp::TriangularSolve {
671                left_side,
672                lower,
673                transpose_a,
674                unit_diagonal,
675            } => {
676                hasher.write_u8(u8::from(left_side));
677                hasher.write_u8(u8::from(lower));
678                hasher.write_u8(u8::from(transpose_a));
679                hasher.write_u8(u8::from(unit_diagonal));
680            }
681            LinalgOp::Cholesky
682            | LinalgOp::Lu
683            | LinalgOp::LuFactor
684            | LinalgOp::LogAbsDetFromLuFactor
685            | LinalgOp::SignDetFromLuFactor
686            | LinalgOp::FullPivLu
687            | LinalgOp::SvdFull
688            | LinalgOp::Solve
689            | LinalgOp::LuFactorSolve
690            | LinalgOp::HouseholderQrFactor
691            | LinalgOp::HouseholderQrFromFactors
692            | LinalgOp::HouseholderQrAppend
693            | LinalgOp::HouseholderQrAppendTangent => {}
694        }
695    }
696
697    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
698        other
699            .as_any()
700            .downcast_ref::<Self>()
701            .is_some_and(|that| self == that)
702    }
703
704    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
705        Arc::new(self.clone())
706    }
707
708    fn as_any(&self) -> &dyn Any {
709        self
710    }
711
712    fn input_count(&self) -> usize {
713        self.op.input_count()
714    }
715
716    fn output_count(&self) -> usize {
717        self.op.output_count()
718    }
719
720    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
721        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
722    }
723
724    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
725        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
726    }
727
728    fn prune_outputs(&self, live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
729        match self.op {
730            LinalgOp::Svd {
731                derivative_eps,
732                driver,
733                ..
734            } if live_outputs == [false, true, false] => {
735                Some(Arc::new(Self::new(LinalgOp::SvdVals {
736                    derivative_eps,
737                    driver,
738                })))
739            }
740            LinalgOp::Eigh {
741                derivative_eps,
742                driver,
743                ..
744            } if live_outputs == [true, false] => Some(Arc::new(Self::new(LinalgOp::EighVals {
745                derivative_eps,
746                driver,
747            }))),
748            LinalgOp::Eig { input_dtype } if live_outputs == [true, false] => {
749                Some(Arc::new(Self::new(LinalgOp::EigVals { input_dtype })))
750            }
751            // `LuFactorSolve` is deliberately never pruned to `Solve`. Traced
752            // AD compiles its source program with this pruning before
753            // differentiating, so a prune would drop the saved factors and
754            // force the adjoint solve to refactor `A`. A primal only program
755            // loses nothing by keeping the op: the fused CPU kernel does the
756            // `Solve` kernel's work, and the factor buffer the plain solve
757            // uses as scratch becomes the (pooled) LU output instead.
758            _ => None,
759        }
760    }
761
762    fn infer_output_meta(
763        &self,
764        ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
765    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
766        let input_dtypes = (0..self.input_count())
767            .map(|input| ctx.input_dtype(input))
768            .collect::<Result<Vec<_>, _>>()?;
769        let input_shapes = (0..self.input_count())
770            .map(|input| ctx.input_shape(input))
771            .collect::<Result<Vec<_>, _>>()?;
772        let metas = match self.op {
773            LinalgOp::Cholesky => {
774                require_matrix_meta("tenferro-linalg.cholesky", input_shapes[0])?;
775                vec![(promote_dtypes(&input_dtypes), input_shapes[0].to_vec())]
776            }
777            LinalgOp::FullPivLuSolve { .. } => {
778                require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[0])?;
779                require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[1])?;
780                vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
781            }
782            LinalgOp::Solve => {
783                require_matrix_meta("tenferro-linalg.solve", input_shapes[0])?;
784                require_matrix_meta("tenferro-linalg.solve", input_shapes[1])?;
785                vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
786            }
787            LinalgOp::LuFactorSolve => {
788                require_matrix_meta("tenferro-linalg.lu_factor_solve", input_shapes[1])?;
789                let mut factors = lu_factor_meta(input_dtypes[0], input_shapes[0])?.into_iter();
790                let (Some(packed_lu), Some(pivots)) = (factors.next(), factors.next()) else {
791                    return Err(Error::Internal(
792                        "lu_factor_solve: lu_factor metadata returned fewer than two outputs"
793                            .into(),
794                    ));
795                };
796                vec![
797                    (promote_dtypes(&input_dtypes), input_shapes[1].to_vec()),
798                    packed_lu,
799                    pivots,
800                ]
801            }
802            LinalgOp::TriangularSolve { .. } => {
803                require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[0])?;
804                require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[1])?;
805                vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
806            }
807            LinalgOp::LuSolvePrepared { .. } => {
808                require_matrix_meta("tenferro-linalg.lu_solve_prepared_lu", input_shapes[0])?;
809                require_matrix_meta("tenferro-linalg.lu_solve_prepared_rhs", input_shapes[3])?;
810                vec![(
811                    promote_dtypes(&[input_dtypes[0], input_dtypes[3]]),
812                    input_shapes[3].to_vec(),
813                )]
814            }
815            LinalgOp::Lu => lu_meta(input_dtypes[0], input_shapes[0])?,
816            LinalgOp::LuFactor => lu_factor_meta(input_dtypes[0], input_shapes[0])?,
817            LinalgOp::SignDetFromLuFactor => {
818                vec![signdet_from_lu_factor_meta(
819                    input_dtypes[0],
820                    input_shapes[0],
821                    input_shapes[1],
822                    input_shapes[2],
823                )?]
824            }
825            LinalgOp::LogAbsDetFromLuFactor => {
826                vec![logabsdet_from_lu_factor_meta(
827                    input_dtypes[0],
828                    input_shapes[0],
829                    input_shapes[1],
830                )?]
831            }
832            LinalgOp::FullPivLu => full_piv_lu_meta(input_dtypes[0], input_shapes[0])?,
833            LinalgOp::Svd { .. } => svd_meta(input_dtypes[0], input_shapes[0])?,
834            LinalgOp::SvdFull => svd_full_meta(input_dtypes[0], input_shapes[0])?,
835            LinalgOp::SvdVals { .. } => {
836                vec![svd_values_meta(input_dtypes[0], input_shapes[0])?]
837            }
838            LinalgOp::Qr { .. } => qr_meta(input_dtypes[0], input_shapes[0])?,
839            LinalgOp::RankRevealingQr { .. } => {
840                rank_revealing_qr_meta(input_dtypes[0], input_shapes[0])?
841            }
842            LinalgOp::HouseholderQrFactor => {
843                householder_qr_factor_meta(input_dtypes[0], input_shapes[0])?
844            }
845            LinalgOp::HouseholderQrFromFactors => {
846                householder_qr_from_factors_meta(&input_dtypes, &input_shapes)?
847            }
848            LinalgOp::HouseholderQrAppend => {
849                householder_qr_append_meta(&input_dtypes, &input_shapes)?
850            }
851            LinalgOp::HouseholderQrR { .. } => vec![householder_qr_r_meta(
852                &input_dtypes,
853                input_shapes[0],
854                input_shapes[1],
855            )?],
856            LinalgOp::HouseholderQrQColumns { start, end, .. } => {
857                vec![householder_qr_q_columns_meta(
858                    &input_dtypes,
859                    input_shapes[0],
860                    input_shapes[1],
861                    start,
862                    end,
863                )?]
864            }
865            LinalgOp::HouseholderQrThinQ { .. } => {
866                vec![householder_qr_thin_q_meta(
867                    &input_dtypes,
868                    input_shapes[0],
869                    input_shapes[1],
870                )?]
871            }
872            LinalgOp::HouseholderQrAppendTangent => {
873                vec![householder_qr_append_tangent_meta(
874                    &input_dtypes,
875                    &input_shapes,
876                )?]
877            }
878            LinalgOp::HouseholderQrSplitTangent { right } => {
879                vec![householder_qr_split_tangent_meta(
880                    &input_dtypes,
881                    &input_shapes,
882                    right,
883                )?]
884            }
885            LinalgOp::Eigh { .. } => eigh_meta(input_dtypes[0], input_shapes[0])?,
886            LinalgOp::EighVals { .. } => vec![eigh_values_meta(input_dtypes[0], input_shapes[0])?],
887            LinalgOp::Eig { input_dtype } => eig_meta(input_dtype, input_shapes[0])?,
888            LinalgOp::EigVals { input_dtype } => {
889                vec![eig_values_meta(input_dtype, input_shapes[0])?]
890            }
891        };
892        Ok(metas)
893    }
894}
895
896fn execute_linalg_extension_reads_on_session<B: BackendSession + ?Sized>(
897    op: &LinalgExtensionOp,
898    inputs: &[TensorRead<'_>],
899    session: &mut B,
900) -> tenferro_tensor::Result<Vec<Tensor>> {
901    if let Some(result) = with_cpu_exec_session(session, |session| {
902        execute_linalg_extension_reads_in_session(op, inputs, session)
903    }) {
904        return result;
905    }
906    #[cfg(feature = "cuda")]
907    if let Some(result) = with_cuda_exec_session(session, |session| {
908        execute_linalg_extension_reads_in_session(op, inputs, session)
909    }) {
910        return result;
911    }
912    Err(Error::unsupported(
913        "linalg_extension",
914        "selected backend session does not expose a linalg execution capability",
915    ))
916}
917
918fn execute_linalg_extension_reads_in_session<S: LinalgBackend>(
919    op: &LinalgExtensionOp,
920    inputs: &[TensorRead<'_>],
921    session: &mut S,
922) -> tenferro_tensor::Result<Vec<Tensor>> {
923    if op.op() == LinalgOp::HouseholderQrAppendTangent {
924        let left = session.to_contiguous_read(inputs[0].clone())?;
925        let right = session.to_contiguous_read(inputs[1].clone())?;
926        return Ok(vec![session.concatenate(&[&left, &right], 1)?]);
927    }
928    if let LinalgOp::HouseholderQrSplitTangent { right } = op.op() {
929        let cotangent_shape = inputs[0].clone().tensor_view().shape().to_vec();
930        let left_shape = inputs[1].clone().tensor_view().shape().to_vec();
931        let right_shape = inputs[2].clone().tensor_view().shape().to_vec();
932        let config =
933            householder_qr_split_config(&cotangent_shape, &left_shape, &right_shape, right)?;
934        let cotangent = session.to_contiguous_read(inputs[0].clone())?;
935        return Ok(vec![session.slice(&cotangent, &config)?]);
936    }
937    // Preserve borrowed layouts through the same provider hooks used by the
938    // concrete read API. In particular, CPU Faer can consume strided views;
939    // packing here would defeat that path before provider dispatch.
940    match op.op() {
941        LinalgOp::Cholesky => return Ok(vec![session.cholesky_read(inputs[0].clone())?]),
942        LinalgOp::Lu => return session.lu_read(inputs[0].clone()),
943        LinalgOp::FullPivLu => return session.full_piv_lu_read(inputs[0].clone()),
944        LinalgOp::Svd {
945            derivative_eps,
946            gauge,
947            driver,
948        } => {
949            return session.svd_with_options_read(
950                inputs[0].clone(),
951                SvdOptions {
952                    derivative_eps,
953                    gauge,
954                    driver,
955                },
956            );
957        }
958        LinalgOp::SvdFull => return session.svd_full_read(inputs[0].clone()),
959        LinalgOp::SvdVals { driver, .. } => {
960            return Ok(vec![
961                session.svd_values_with_driver_read(inputs[0].clone(), driver)?
962            ]);
963        }
964        LinalgOp::Qr { gauge } => {
965            return session.qr_with_options_read(inputs[0].clone(), QrOptions { gauge });
966        }
967        LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
968            return session.rank_revealing_qr_read(
969                inputs[0].clone(),
970                RankRevealingQrOptions { gauge, rtol, atol },
971            );
972        }
973        LinalgOp::Eigh {
974            derivative_eps,
975            gauge,
976            driver,
977        } => {
978            return session.eigh_with_options_read(
979                inputs[0].clone(),
980                EighOptions {
981                    derivative_eps,
982                    gauge,
983                    driver,
984                },
985            );
986        }
987        LinalgOp::EighVals { driver, .. } => {
988            return Ok(vec![
989                session.eigh_values_with_driver_read(inputs[0].clone(), driver)?
990            ]);
991        }
992        LinalgOp::Eig { .. } => return session.eig_read(inputs[0].clone()),
993        LinalgOp::EigVals { .. } => return Ok(vec![session.eig_values_read(inputs[0].clone())?]),
994        LinalgOp::Solve => match session.solve_read(inputs[0].clone(), inputs[1].clone()) {
995            Ok(output) => return Ok(vec![output]),
996            Err(error) if error.kind() == ErrorKind::Unsupported => {}
997            Err(error) => return Err(error),
998        },
999        _ => {}
1000    }
1001    if let LinalgOp::TriangularSolve {
1002        left_side,
1003        lower,
1004        transpose_a,
1005        unit_diagonal,
1006    } = op.op()
1007    {
1008        match session.triangular_solve_read(
1009            inputs[0].clone(),
1010            inputs[1].clone(),
1011            left_side,
1012            lower,
1013            transpose_a,
1014            unit_diagonal,
1015        ) {
1016            Ok(output) => return Ok(vec![output]),
1017            Err(error) if error.kind() == ErrorKind::Unsupported => {}
1018            Err(error) => return Err(error),
1019        }
1020    }
1021
1022    // The owned-only hooks need materialized views, not copies of inputs that
1023    // are already owned tensors (notably the saved matrix and LU in backward).
1024    let materialized_inputs = inputs
1025        .iter()
1026        .filter(|input| input.as_tensor().is_none())
1027        .cloned()
1028        .map(|input| session.to_contiguous_read(input))
1029        .collect::<tenferro_tensor::Result<Vec<_>>>()?;
1030    let mut views = materialized_inputs.iter();
1031    // INVARIANT: exactly one materialized tensor was produced per View, in
1032    // input order; every input therefore contributes exactly one reference.
1033    let input_refs: Vec<&Tensor> = inputs
1034        .iter()
1035        .filter_map(|input| input.as_tensor().or_else(|| views.next()))
1036        .collect();
1037    execute_linalg(op.op(), &input_refs, session)
1038}
1039
1040fn linalg_session_supported<B: tenferro_tensor::TensorBackend + 'static>(
1041    #[cfg_attr(not(feature = "cuda"), allow(unused_variables))] op: &LinalgExtensionOp,
1042) -> bool {
1043    // The `supports_session` contract (capability.rs) requires that an op is
1044    // admitted to a scheduler session only when the session executor genuinely
1045    // executes it without returning `Unsupported`. Admission is exactly
1046    // per-op/per-backend so `apply_eager` keeps the native prepared path for
1047    // every op the session can actually run (issue #1665).
1048    let type_id = std::any::TypeId::of::<B>();
1049    if type_id == std::any::TypeId::of::<tenferro_cpu::CpuBackend>() {
1050        // Every CPU linalg kernel, including full-matrices SVD, runs in-session
1051        // on both the faer and the BLAS provider, so the type-only seam can
1052        // admit the whole family without inspecting the provider kind.
1053        return true;
1054    }
1055    #[cfg(feature = "cuda")]
1056    {
1057        if type_id == std::any::TypeId::of::<tenferro_gpu::cuda::CudaBackend>() {
1058            return match op.op() {
1059                // Complete-pivoting LU and general eig have no CUDA kernels.
1060                LinalgOp::FullPivLu | LinalgOp::FullPivLuSolve { .. } => false,
1061                LinalgOp::Eig { .. } | LinalgOp::EigVals { .. } => false,
1062                // Plain partial-pivot solve runs in-session via cuSOLVER
1063                // getrf plus prepared pivot/triangular solves
1064                // (`gpu/linalg.rs::solve` = lu_factor + lu_solve_prepared, no
1065                // Unsupported path for F32/F64/C32/C64), so it is admitted.
1066                LinalgOp::Solve => true,
1067                // The fused solve runs the trait default on CUDA: getrf then
1068                // the plain prepared solve, both admitted above.
1069                LinalgOp::LuFactorSolve => true,
1070                // Conjugate-only prepared LU solve is unsupported on CUDA.
1071                LinalgOp::LuSolvePrepared {
1072                    transpose_a: false,
1073                    conjugate_a: true,
1074                } => false,
1075                _ => true,
1076            };
1077        }
1078    }
1079    false
1080}
1081
1082fn execute_linalg_extension_in_session(
1083    op: &LinalgExtensionOp,
1084    session: &mut dyn BackendSession,
1085    _extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
1086    inputs: &[TensorRead<'_>],
1087) -> tenferro_tensor::Result<Vec<Tensor>> {
1088    // Reuse the existing session executor that the eager and scheduler paths
1089    // already share; it downcasts the borrowed session to the CPU/CUDA exec
1090    // session and runs the same forward kernel for every LinalgOp.
1091    execute_linalg_extension_reads_on_session(op, inputs, session)
1092}
1093
1094define_extension_runtime! {
1095    runtime = LinalgRuntime,
1096    family_id = LINALG_EXTENSION_FAMILY_ID,
1097    op_type = LinalgExtensionOp,
1098    execute_in_session = execute_linalg_extension_in_session,
1099    session_supported = linalg_session_supported,
1100    backend_bound = TensorBackend,
1101}
1102
1103fn execute_linalg<B: LinalgBackend>(
1104    op: LinalgOp,
1105    inputs: &[&Tensor],
1106    backend: &mut B,
1107) -> tenferro_tensor::Result<Vec<Tensor>> {
1108    match op {
1109        LinalgOp::Cholesky => Ok(vec![backend.cholesky(inputs[0])?]),
1110        LinalgOp::Lu => backend.lu(inputs[0]),
1111        LinalgOp::LuFactor => backend.lu_factor(inputs[0]),
1112        LinalgOp::SignDetFromLuFactor => {
1113            Ok(vec![signdet_from_lu_factor(inputs[1], inputs[2], backend)?])
1114        }
1115        LinalgOp::LogAbsDetFromLuFactor => Ok(vec![logabsdet_from_lu_factor(inputs[1], backend)?]),
1116        LinalgOp::LuSolvePrepared {
1117            transpose_a,
1118            conjugate_a,
1119        } => Ok(vec![backend.lu_solve_prepared(
1120            inputs[0],
1121            inputs[1],
1122            inputs[2],
1123            inputs[3],
1124            transpose_a,
1125            conjugate_a,
1126        )?]),
1127        LinalgOp::FullPivLu => backend.full_piv_lu(inputs[0]),
1128        LinalgOp::FullPivLuSolve { transpose_a } => Ok(vec![backend.full_piv_lu_solve(
1129            inputs[0],
1130            inputs[1],
1131            transpose_a,
1132        )?]),
1133        LinalgOp::Solve => Ok(vec![backend.solve(inputs[0], inputs[1])?]),
1134        LinalgOp::LuFactorSolve => backend.lu_factor_solve(inputs[0], inputs[1]),
1135        LinalgOp::Svd {
1136            derivative_eps,
1137            gauge,
1138            driver,
1139        } => backend.svd_with_options(
1140            inputs[0],
1141            SvdOptions {
1142                derivative_eps,
1143                gauge,
1144                driver,
1145            },
1146        ),
1147        LinalgOp::SvdFull => backend.svd_full(inputs[0]),
1148        LinalgOp::SvdVals { driver, .. } => {
1149            Ok(vec![backend.svd_values_with_driver(inputs[0], driver)?])
1150        }
1151        LinalgOp::Qr { gauge } => backend.qr_with_options(inputs[0], QrOptions { gauge }),
1152        LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
1153            backend.rank_revealing_qr(inputs[0], RankRevealingQrOptions { gauge, rtol, atol })
1154        }
1155        LinalgOp::HouseholderQrFactor => {
1156            let state = backend.householder_qr(inputs[0])?;
1157            Ok(vec![state.packed, state.coeff])
1158        }
1159        LinalgOp::HouseholderQrFromFactors => {
1160            let state = backend.householder_qr_from_factors(inputs[0], inputs[1])?;
1161            Ok(vec![state.packed, state.coeff])
1162        }
1163        LinalgOp::HouseholderQrAppend => {
1164            let state = backend.householder_qr_append(inputs[0], inputs[1], inputs[2])?;
1165            Ok(vec![state.packed, state.coeff])
1166        }
1167        LinalgOp::HouseholderQrR { gauge } => Ok(vec![backend.householder_qr_r(
1168            inputs[0],
1169            inputs[1],
1170            QrOptions { gauge },
1171        )?]),
1172        LinalgOp::HouseholderQrQColumns { start, end, gauge } => Ok(vec![backend
1173            .householder_qr_q_columns(inputs[0], inputs[1], start..end, QrOptions { gauge })?]),
1174        LinalgOp::HouseholderQrThinQ { gauge } => {
1175            let end = inputs[1].shape().first().copied().ok_or_else(|| {
1176                Error::rank_mismatch("tenferro-linalg.householder_qr_thin_q", 1, 0)
1177            })?;
1178            Ok(vec![backend.householder_qr_q_columns(
1179                inputs[0],
1180                inputs[1],
1181                0..end,
1182                QrOptions { gauge },
1183            )?])
1184        }
1185        LinalgOp::HouseholderQrAppendTangent => {
1186            Ok(vec![backend.concatenate(&[inputs[0], inputs[1]], 1)?])
1187        }
1188        LinalgOp::HouseholderQrSplitTangent { right } => {
1189            let config = householder_qr_split_config(
1190                inputs[0].shape(),
1191                inputs[1].shape(),
1192                inputs[2].shape(),
1193                right,
1194            )?;
1195            Ok(vec![backend.slice(inputs[0], &config)?])
1196        }
1197        LinalgOp::Eigh {
1198            derivative_eps,
1199            gauge,
1200            driver,
1201        } => backend.eigh_with_options(
1202            inputs[0],
1203            EighOptions {
1204                derivative_eps,
1205                gauge,
1206                driver,
1207            },
1208        ),
1209        LinalgOp::EighVals { driver, .. } => {
1210            Ok(vec![backend.eigh_values_with_driver(inputs[0], driver)?])
1211        }
1212        LinalgOp::Eig { .. } => backend.eig(inputs[0]),
1213        LinalgOp::EigVals { .. } => Ok(vec![backend.eig_values(inputs[0])?]),
1214        LinalgOp::TriangularSolve {
1215            left_side,
1216            lower,
1217            transpose_a,
1218            unit_diagonal,
1219        } => Ok(vec![backend.triangular_solve(
1220            inputs[0],
1221            inputs[1],
1222            left_side,
1223            lower,
1224            transpose_a,
1225            unit_diagonal,
1226        )?]),
1227    }
1228}
1229
1230/// Sign (or complex phase) of the determinant from an LU factorization.
1231///
1232/// The sign is built from the per-pivot signs rather than from the determinant
1233/// product, so it is magnitude-independent: a determinant whose product would
1234/// underflow or overflow still reports its mathematical sign. `Sign` maps an
1235/// exactly zero pivot to zero, so finite singular LU factors report zero sign
1236/// for both real and complex inputs, matching the existing eager composite.
1237fn signdet_from_lu_factor<B: LinalgBackend + ?Sized>(
1238    packed_lu: &Tensor,
1239    parity: &Tensor,
1240    backend: &mut B,
1241) -> tenferro_tensor::Result<Tensor> {
1242    let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
1243    let sign_diag = backend.sign_read(TensorRead::from_tensor(&diag))?;
1244    let sign_u = backend.reduce_prod_read(TensorRead::from_tensor(&sign_diag), &[0])?;
1245    backend.mul_read(
1246        TensorRead::from_tensor(parity),
1247        TensorRead::from_tensor(&sign_u),
1248    )
1249}
1250
1251fn logabsdet_from_lu_factor<B: LinalgBackend + ?Sized>(
1252    packed_lu: &Tensor,
1253    backend: &mut B,
1254) -> tenferro_tensor::Result<Tensor> {
1255    let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
1256    let abs = backend.abs_read(TensorRead::from_tensor(&diag))?;
1257    let log = backend.log_read(TensorRead::from_tensor(&abs))?;
1258    backend.reduce_sum_read(TensorRead::from_tensor(&log), &[0])
1259}
1260
1261pub(crate) fn apply_svd_gauge(
1262    gauge: SvdGauge,
1263    outputs: &mut [Tensor],
1264) -> tenferro_tensor::Result<()> {
1265    match gauge {
1266        SvdGauge::Raw => Ok(()),
1267        SvdGauge::CanonicalPivot => apply_canonical_pivot_svd_gauge(outputs),
1268    }
1269}
1270
1271fn apply_canonical_pivot_svd_gauge(outputs: &mut [Tensor]) -> tenferro_tensor::Result<()> {
1272    if outputs.len() != 3 {
1273        return Err(Error::invalid_argument(
1274            "tenferro-linalg.svd",
1275            "outputs",
1276            format!(
1277                "canonical SVD gauge expected three outputs, got {}",
1278                outputs.len()
1279            ),
1280        ));
1281    }
1282
1283    let (u_slice, rest) = outputs.split_at_mut(1);
1284    let (singular_slice, vt_slice) = rest.split_at_mut(1);
1285    let u = &mut u_slice[0];
1286    let singular_values = &singular_slice[0];
1287    let vt = &mut vt_slice[0];
1288    let u_shape = u.shape().to_vec();
1289    let s_shape = singular_values.shape().to_vec();
1290    let vt_shape = vt.shape().to_vec();
1291    if u_shape.len() < 2 || vt_shape.len() < 2 || s_shape.is_empty() {
1292        return Err(Error::invalid_argument(
1293            "tenferro-linalg.svd",
1294            "outputs",
1295            format!(
1296                "canonical SVD gauge expected U rank >= 2, S rank >= 1, VT rank >= 2; got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
1297            ),
1298        ));
1299    }
1300
1301    let m = u_shape[0];
1302    let k = u_shape[1];
1303    let n = vt_shape[1];
1304    if s_shape[0] != k
1305        || vt_shape[0] != k
1306        || u_shape[2..] != vt_shape[2..]
1307        || s_shape[1..] != u_shape[2..]
1308    {
1309        return Err(Error::invalid_argument(
1310            "tenferro-linalg.svd",
1311            "outputs",
1312            format!(
1313                "canonical SVD gauge expected compatible compact SVD shapes, got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
1314            ),
1315        ));
1316    }
1317    let layout = canonical_svd_gauge_layout(m, k, n, &u_shape[2..])?;
1318
1319    match (u.dtype(), vt.dtype()) {
1320        (tenferro_tensor::DType::F64, tenferro_tensor::DType::F64) => {
1321            let (u, vt) = svd_gauge_pair_mut::<f64>(u, vt)?;
1322            canonicalize_svd_gauge_f64(u.host_data_mut()?, vt.host_data_mut()?, layout)
1323        }
1324        (tenferro_tensor::DType::F32, tenferro_tensor::DType::F32) => {
1325            let (u, vt) = svd_gauge_pair_mut::<f32>(u, vt)?;
1326            canonicalize_svd_gauge_f32(u.host_data_mut()?, vt.host_data_mut()?, layout)
1327        }
1328        (tenferro_tensor::DType::C64, tenferro_tensor::DType::C64) => {
1329            let (u, vt) = svd_gauge_pair_mut::<Complex64>(u, vt)?;
1330            canonicalize_svd_gauge_c64(u.host_data_mut()?, vt.host_data_mut()?, layout)
1331        }
1332        (tenferro_tensor::DType::C32, tenferro_tensor::DType::C32) => {
1333            let (u, vt) = svd_gauge_pair_mut::<Complex32>(u, vt)?;
1334            canonicalize_svd_gauge_c32(u.host_data_mut()?, vt.host_data_mut()?, layout)
1335        }
1336        (u_dtype, vt_dtype) => Err(Error::dtype_mismatch(
1337            "tenferro-linalg.svd",
1338            u_dtype,
1339            vt_dtype,
1340        )),
1341    }
1342}
1343
1344/// The typed pair behind a same-dtype pair of mutable tensors, or this module's refusal.
1345fn svd_gauge_pair_mut<'a, T: tenferro_tensor::TensorScalar>(
1346    u: &'a mut Tensor,
1347    vt: &'a mut Tensor,
1348) -> tenferro_tensor::Result<(
1349    &'a mut tenferro_tensor::TypedTensor<T>,
1350    &'a mut tenferro_tensor::TypedTensor<T>,
1351)> {
1352    let (u_dtype, vt_dtype) = (u.dtype(), vt.dtype());
1353    let u_t = u
1354        .as_typed_mut::<T>()
1355        .ok_or_else(|| Error::dtype_mismatch("tenferro-linalg.svd", u_dtype, vt_dtype))?;
1356    let vt_t = vt
1357        .as_typed_mut::<T>()
1358        .ok_or_else(|| Error::dtype_mismatch("tenferro-linalg.svd", u_dtype, vt_dtype))?;
1359    Ok((u_t, vt_t))
1360}
1361
1362#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1363struct CanonicalSvdGaugeLayout {
1364    m: usize,
1365    k: usize,
1366    batch_count: usize,
1367    u_batch_len: usize,
1368    vt_batch_len: usize,
1369    u_len: usize,
1370    vt_len: usize,
1371}
1372
1373impl CanonicalSvdGaugeLayout {
1374    fn validate_storage(self, u_len: usize, vt_len: usize) -> tenferro_tensor::Result<()> {
1375        if u_len != self.u_len {
1376            return Err(Error::invalid_argument(
1377                "tenferro-linalg.svd",
1378                "U storage",
1379                format!(
1380                    "canonical SVD gauge expected U storage length {}, got {u_len}",
1381                    self.u_len
1382                ),
1383            ));
1384        }
1385        if vt_len != self.vt_len {
1386            return Err(Error::invalid_argument(
1387                "tenferro-linalg.svd",
1388                "VT storage",
1389                format!(
1390                    "canonical SVD gauge expected VT storage length {}, got {vt_len}",
1391                    self.vt_len
1392                ),
1393            ));
1394        }
1395        Ok(())
1396    }
1397}
1398
1399fn canonical_svd_gauge_layout(
1400    m: usize,
1401    k: usize,
1402    n: usize,
1403    batch_shape: &[usize],
1404) -> tenferro_tensor::Result<CanonicalSvdGaugeLayout> {
1405    let batch_count = tenferro_tensor::validate::checked_shape_product(
1406        "tenferro-linalg.svd",
1407        "canonical SVD batch",
1408        batch_shape,
1409    )?;
1410    let u_batch_len = tenferro_tensor::validate::checked_shape_product(
1411        "tenferro-linalg.svd",
1412        "canonical SVD U batch",
1413        &[m, k],
1414    )?;
1415    let vt_batch_len = tenferro_tensor::validate::checked_shape_product(
1416        "tenferro-linalg.svd",
1417        "canonical SVD VT batch",
1418        &[k, n],
1419    )?;
1420    let u_len = tenferro_tensor::validate::checked_shape_product(
1421        "tenferro-linalg.svd",
1422        "canonical SVD U storage",
1423        &[u_batch_len, batch_count],
1424    )?;
1425    let vt_len = tenferro_tensor::validate::checked_shape_product(
1426        "tenferro-linalg.svd",
1427        "canonical SVD VT storage",
1428        &[vt_batch_len, batch_count],
1429    )?;
1430    Ok(CanonicalSvdGaugeLayout {
1431        m,
1432        k,
1433        batch_count,
1434        u_batch_len,
1435        vt_batch_len,
1436        u_len,
1437        vt_len,
1438    })
1439}
1440
1441fn canonicalize_svd_gauge_f64(
1442    u: &mut [f64],
1443    vt: &mut [f64],
1444    layout: CanonicalSvdGaugeLayout,
1445) -> tenferro_tensor::Result<()> {
1446    layout.validate_storage(u.len(), vt.len())?;
1447    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1448        return Ok(());
1449    }
1450    for (u_batch, vt_batch) in u
1451        .chunks_exact_mut(layout.u_batch_len)
1452        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1453    {
1454        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1455            let pivot = max_abs_pivot_f64(u_column);
1456            let pivot_value = u_column[pivot];
1457            if pivot_value < 0.0 {
1458                for value in u_column {
1459                    *value = -*value;
1460                }
1461                for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1462                    vt_column[col] = -vt_column[col];
1463                }
1464            }
1465        }
1466    }
1467    Ok(())
1468}
1469
1470fn canonicalize_svd_gauge_f32(
1471    u: &mut [f32],
1472    vt: &mut [f32],
1473    layout: CanonicalSvdGaugeLayout,
1474) -> tenferro_tensor::Result<()> {
1475    layout.validate_storage(u.len(), vt.len())?;
1476    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1477        return Ok(());
1478    }
1479    for (u_batch, vt_batch) in u
1480        .chunks_exact_mut(layout.u_batch_len)
1481        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1482    {
1483        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1484            let pivot = max_abs_pivot_f32(u_column);
1485            let pivot_value = u_column[pivot];
1486            if pivot_value < 0.0 {
1487                for value in u_column {
1488                    *value = -*value;
1489                }
1490                for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1491                    vt_column[col] = -vt_column[col];
1492                }
1493            }
1494        }
1495    }
1496    Ok(())
1497}
1498
1499fn canonicalize_svd_gauge_c64(
1500    u: &mut [Complex64],
1501    vt: &mut [Complex64],
1502    layout: CanonicalSvdGaugeLayout,
1503) -> tenferro_tensor::Result<()> {
1504    layout.validate_storage(u.len(), vt.len())?;
1505    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1506        return Ok(());
1507    }
1508    for (u_batch, vt_batch) in u
1509        .chunks_exact_mut(layout.u_batch_len)
1510        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1511    {
1512        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1513            let pivot = max_abs_pivot_c64(u_column);
1514            let pivot_value = u_column[pivot];
1515            let pivot_norm = pivot_value.norm();
1516            if pivot_norm == 0.0 {
1517                continue;
1518            }
1519            let phase = pivot_value.conj() / pivot_norm;
1520            let vt_phase = phase.conj();
1521            for value in u_column {
1522                *value *= phase;
1523            }
1524            for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1525                vt_column[col] *= vt_phase;
1526            }
1527        }
1528    }
1529    Ok(())
1530}
1531
1532fn canonicalize_svd_gauge_c32(
1533    u: &mut [Complex32],
1534    vt: &mut [Complex32],
1535    layout: CanonicalSvdGaugeLayout,
1536) -> tenferro_tensor::Result<()> {
1537    layout.validate_storage(u.len(), vt.len())?;
1538    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1539        return Ok(());
1540    }
1541    for (u_batch, vt_batch) in u
1542        .chunks_exact_mut(layout.u_batch_len)
1543        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1544    {
1545        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1546            let pivot = max_abs_pivot_c32(u_column);
1547            let pivot_value = u_column[pivot];
1548            let pivot_norm = pivot_value.norm();
1549            if pivot_norm == 0.0 {
1550                continue;
1551            }
1552            let phase = pivot_value.conj() / pivot_norm;
1553            let vt_phase = phase.conj();
1554            for value in u_column {
1555                *value *= phase;
1556            }
1557            for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1558                vt_column[col] *= vt_phase;
1559            }
1560        }
1561    }
1562    Ok(())
1563}
1564
1565fn max_abs_pivot_f64(u_column: &[f64]) -> usize {
1566    let mut pivot = 0;
1567    let mut pivot_abs = u_column[0].abs();
1568    for (row, value) in u_column.iter().enumerate().skip(1) {
1569        let candidate_abs = value.abs();
1570        if candidate_abs > pivot_abs {
1571            pivot = row;
1572            pivot_abs = candidate_abs;
1573        }
1574    }
1575    pivot
1576}
1577
1578fn max_abs_pivot_f32(u_column: &[f32]) -> usize {
1579    let mut pivot = 0;
1580    let mut pivot_abs = u_column[0].abs();
1581    for (row, value) in u_column.iter().enumerate().skip(1) {
1582        let candidate_abs = value.abs();
1583        if candidate_abs > pivot_abs {
1584            pivot = row;
1585            pivot_abs = candidate_abs;
1586        }
1587    }
1588    pivot
1589}
1590
1591fn max_abs_pivot_c64(u_column: &[Complex64]) -> usize {
1592    let mut pivot = 0;
1593    let mut pivot_abs = u_column[0].norm_sqr();
1594    for (row, value) in u_column.iter().enumerate().skip(1) {
1595        let candidate_abs = value.norm_sqr();
1596        if candidate_abs > pivot_abs {
1597            pivot = row;
1598            pivot_abs = candidate_abs;
1599        }
1600    }
1601    pivot
1602}
1603
1604fn max_abs_pivot_c32(u_column: &[Complex32]) -> usize {
1605    let mut pivot = 0;
1606    let mut pivot_abs = u_column[0].norm_sqr();
1607    for (row, value) in u_column.iter().enumerate().skip(1) {
1608        let candidate_abs = value.norm_sqr();
1609        if candidate_abs > pivot_abs {
1610            pivot = row;
1611            pivot_abs = candidate_abs;
1612        }
1613    }
1614    pivot
1615}
1616
1617fn require_matrix_meta(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1618    if shape.len() < 2 {
1619        return Err(Error::rank_mismatch(op, 2, shape.len()));
1620    }
1621    Ok(())
1622}
1623
1624fn matrix_meta_parts<'a>(
1625    op: &'static str,
1626    shape: &'a [SymDim],
1627) -> tenferro_tensor::Result<(SymDim, SymDim, &'a [SymDim])> {
1628    require_matrix_meta(op, shape)?;
1629    Ok((shape[0].clone(), shape[1].clone(), &shape[2..]))
1630}
1631
1632fn lu_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1633    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu", shape)?;
1634    let k = m.clone().min(n.clone());
1635    Ok(vec![
1636        (dtype, matrix_shape(m.clone(), m, batch)),
1637        (dtype, matrix_shape(shape[0].clone(), k.clone(), batch)),
1638        (dtype, matrix_shape(k, n, batch)),
1639        (dtype, batch.to_vec()),
1640    ])
1641}
1642
1643fn lu_factor_meta(
1644    dtype: DType,
1645    shape: &[SymDim],
1646) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1647    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu_factor", shape)?;
1648    let k = m.min(n);
1649    Ok(vec![
1650        (dtype, shape.to_vec()),
1651        (DType::I32, vector_shape(k, batch)),
1652        (dtype, batch.to_vec()),
1653    ])
1654}
1655
1656fn signdet_from_lu_factor_meta(
1657    input_dtype: DType,
1658    input_shape: &[SymDim],
1659    packed_shape: &[SymDim],
1660    parity_shape: &[SymDim],
1661) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1662    let (_, _, batch) = matrix_meta_parts("tenferro-linalg.signdet_from_lu_factor", input_shape)?;
1663    require_matrix_meta(
1664        "tenferro-linalg.signdet_from_lu_factor_packed",
1665        packed_shape,
1666    )?;
1667    if parity_shape.len() != batch.len() {
1668        return Err(Error::rank_mismatch(
1669            "tenferro-linalg.signdet_from_lu_factor_parity",
1670            batch.len(),
1671            parity_shape.len(),
1672        ));
1673    }
1674    Ok((input_dtype, batch.to_vec()))
1675}
1676
1677fn logabsdet_from_lu_factor_meta(
1678    input_dtype: DType,
1679    input_shape: &[SymDim],
1680    packed_shape: &[SymDim],
1681) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1682    let (_, _, batch) = matrix_meta_parts("tenferro-linalg.logabsdet_from_lu_factor", input_shape)?;
1683    require_matrix_meta(
1684        "tenferro-linalg.logabsdet_from_lu_factor_packed",
1685        packed_shape,
1686    )?;
1687    Ok((singular_values_dtype(input_dtype), batch.to_vec()))
1688}
1689
1690fn full_piv_lu_meta(
1691    dtype: DType,
1692    shape: &[SymDim],
1693) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1694    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.full_piv_lu", shape)?;
1695    Ok(vec![
1696        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1697        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1698        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1699        (dtype, matrix_shape(n.clone(), n, batch)),
1700        (singular_values_dtype(dtype), batch.to_vec()),
1701    ])
1702}
1703
1704fn svd_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1705    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd", shape)?;
1706    let k = m.clone().min(n.clone());
1707    Ok(vec![
1708        (dtype, matrix_shape(m, k.clone(), batch)),
1709        (singular_values_dtype(dtype), vector_shape(k.clone(), batch)),
1710        (dtype, matrix_shape(k, n, batch)),
1711    ])
1712}
1713
1714fn svd_full_meta(
1715    dtype: DType,
1716    shape: &[SymDim],
1717) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1718    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_full", shape)?;
1719    let k = m.clone().min(n.clone());
1720    Ok(vec![
1721        (dtype, matrix_shape(m.clone(), m, batch)),
1722        (singular_values_dtype(dtype), vector_shape(k, batch)),
1723        (dtype, matrix_shape(n.clone(), n, batch)),
1724    ])
1725}
1726
1727fn svd_values_meta(
1728    dtype: DType,
1729    shape: &[SymDim],
1730) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1731    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_values", shape)?;
1732    let k = m.min(n);
1733    Ok((singular_values_dtype(dtype), vector_shape(k, batch)))
1734}
1735
1736fn qr_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1737    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.qr", shape)?;
1738    let k = m.clone().min(n.clone());
1739    Ok(vec![
1740        (dtype, matrix_shape(m, k.clone(), batch)),
1741        (dtype, matrix_shape(k, n, batch)),
1742    ])
1743}
1744
1745fn rank_revealing_qr_meta(
1746    dtype: DType,
1747    shape: &[SymDim],
1748) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1749    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.rank_revealing_qr", shape)?;
1750    let k = m.clone().min(n.clone());
1751    Ok(vec![
1752        (dtype, matrix_shape(m, k.clone(), batch)),
1753        (dtype, matrix_shape(k, n.clone(), batch)),
1754        (DType::I64, vector_shape(n, batch)),
1755        (DType::I64, batch.to_vec()),
1756    ])
1757}
1758
1759fn require_householder_rank2(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1760    if shape.len() != 2 {
1761        return Err(Error::rank_mismatch(op, 2, shape.len()));
1762    }
1763    Ok(())
1764}
1765
1766fn householder_qr_factor_meta(
1767    dtype: DType,
1768    shape: &[SymDim],
1769) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1770    require_householder_rank2("tenferro-linalg.householder_qr", shape)?;
1771    let k = shape[0].clone().min(shape[1].clone());
1772    Ok(vec![(dtype, shape.to_vec()), (dtype, vec![k])])
1773}
1774
1775fn householder_qr_from_factors_meta(
1776    dtypes: &[DType],
1777    shapes: &[&[SymDim]],
1778) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1779    const OP: &str = "tenferro-linalg.householder_qr_from_factors";
1780    require_householder_rank2(OP, shapes[0])?;
1781    require_householder_rank2(OP, shapes[1])?;
1782    if dtypes[0] != dtypes[1] {
1783        return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1784    }
1785    require_static_extent_equal(OP, "q.cols/r.rows", &shapes[0][1], &shapes[1][0])?;
1786    if let (Some(q_cols), Some(q_rows), Some(r_cols)) = (
1787        shapes[0][1].constant_value(),
1788        shapes[0][0].constant_value(),
1789        shapes[1][1].constant_value(),
1790    ) && q_cols > q_rows.min(r_cols)
1791    {
1792        return Err(Error::invalid_argument(
1793            OP,
1794            "shape",
1795            "Q column count exceeds min(Q rows, R columns)",
1796        ));
1797    }
1798    let m = shapes[0][0].clone();
1799    let n = shapes[1][1].clone();
1800    let k = m.clone().min(n.clone());
1801    Ok(vec![(dtypes[0], vec![m, n]), (dtypes[0], vec![k])])
1802}
1803
1804fn householder_qr_append_meta(
1805    dtypes: &[DType],
1806    shapes: &[&[SymDim]],
1807) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1808    const OP: &str = "tenferro-linalg.householder_qr_append";
1809    require_householder_state_meta(OP, dtypes, shapes[0], shapes[1])?;
1810    require_householder_rank2(OP, shapes[2])?;
1811    if dtypes[0] != dtypes[2] {
1812        return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[2]));
1813    }
1814    require_static_extent_equal(OP, "rows", &shapes[0][0], &shapes[2][0])?;
1815    let m = shapes[0][0].clone();
1816    let width = shapes[0][1].clone() + shapes[2][1].clone();
1817    let k = m.clone().min(width.clone());
1818    Ok(vec![(dtypes[0], vec![m, width]), (dtypes[0], vec![k])])
1819}
1820
1821fn require_static_extent_equal(
1822    op: &'static str,
1823    field: &'static str,
1824    lhs: &SymDim,
1825    rhs: &SymDim,
1826) -> tenferro_tensor::Result<()> {
1827    if let (Some(lhs), Some(rhs)) = (lhs.constant_value(), rhs.constant_value())
1828        && lhs != rhs
1829    {
1830        return Err(Error::invalid_argument(
1831            op,
1832            field,
1833            format!("expected equal extents, got {lhs} and {rhs}"),
1834        ));
1835    }
1836    Ok(())
1837}
1838
1839fn require_householder_state_meta(
1840    op: &'static str,
1841    dtypes: &[DType],
1842    packed: &[SymDim],
1843    coeff: &[SymDim],
1844) -> tenferro_tensor::Result<()> {
1845    require_householder_rank2(op, packed)?;
1846    if coeff.len() != 1 {
1847        return Err(Error::rank_mismatch(op, 1, coeff.len()));
1848    }
1849    if dtypes[0] != dtypes[1] {
1850        return Err(Error::dtype_mismatch(op, dtypes[0], dtypes[1]));
1851    }
1852    let expected = packed[0].clone().min(packed[1].clone());
1853    require_static_extent_equal(op, "coeff", &coeff[0], &expected)
1854}
1855
1856fn householder_qr_r_meta(
1857    dtypes: &[DType],
1858    packed: &[SymDim],
1859    coeff: &[SymDim],
1860) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1861    require_householder_state_meta("tenferro-linalg.householder_qr_r", dtypes, packed, coeff)?;
1862    Ok((dtypes[0], vec![coeff[0].clone(), packed[1].clone()]))
1863}
1864
1865fn householder_qr_q_columns_meta(
1866    dtypes: &[DType],
1867    packed: &[SymDim],
1868    coeff: &[SymDim],
1869    start: usize,
1870    end: usize,
1871) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1872    require_householder_state_meta(
1873        "tenferro-linalg.householder_qr_q_columns",
1874        dtypes,
1875        packed,
1876        coeff,
1877    )?;
1878    if start > end {
1879        return Err(Error::invalid_argument(
1880            "tenferro-linalg.householder_qr_q_columns",
1881            "range",
1882            format!("invalid Q-column range {start}..{end}"),
1883        ));
1884    }
1885    // The reachable width is full Q, not thin Q: columns `k..m` span the
1886    // orthogonal complement of the input's column space.
1887    if packed[0].constant_value().is_some_and(|rows| end > rows) {
1888        return Err(Error::invalid_argument(
1889            "tenferro-linalg.householder_qr_q_columns",
1890            "range",
1891            format!("Q-column range {start}..{end} exceeds full-Q width"),
1892        ));
1893    }
1894    Ok((
1895        dtypes[0],
1896        vec![packed[0].clone(), SymDim::from(end - start)],
1897    ))
1898}
1899
1900fn householder_qr_thin_q_meta(
1901    dtypes: &[DType],
1902    packed: &[SymDim],
1903    coeff: &[SymDim],
1904) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1905    require_householder_state_meta(
1906        "tenferro-linalg.householder_qr_thin_q",
1907        dtypes,
1908        packed,
1909        coeff,
1910    )?;
1911    Ok((dtypes[0], vec![packed[0].clone(), coeff[0].clone()]))
1912}
1913
1914fn householder_qr_split_config(
1915    cotangent: &[usize],
1916    left: &[usize],
1917    right_shape: &[usize],
1918    take_right: bool,
1919) -> tenferro_tensor::Result<tenferro_tensor::SliceConfig> {
1920    const OP: &str = "tenferro-linalg.householder_qr_split_tangent";
1921    for shape in [cotangent, left, right_shape] {
1922        if shape.len() != 2 {
1923            return Err(Error::rank_mismatch(OP, 2, shape.len()));
1924        }
1925    }
1926    let total_width = left[1]
1927        .checked_add(right_shape[1])
1928        .ok_or_else(|| Error::invalid_argument(OP, "shape", "column range overflow"))?;
1929    if cotangent[0] != left[0] || cotangent[0] != right_shape[0] || cotangent[1] != total_width {
1930        return Err(Error::invalid_argument(
1931            OP,
1932            "shape",
1933            "cotangent shape does not match appended factors",
1934        ));
1935    }
1936    let selected = if take_right { right_shape } else { left };
1937    let start = if take_right { left[1] } else { 0 };
1938    let end = start
1939        .checked_add(selected[1])
1940        .ok_or_else(|| Error::invalid_argument(OP, "shape", "column range overflow"))?;
1941    Ok(tenferro_tensor::SliceConfig {
1942        starts: vec![0, start],
1943        limits: vec![selected[0], end],
1944        strides: vec![1, 1],
1945    })
1946}
1947
1948fn householder_qr_append_tangent_meta(
1949    dtypes: &[DType],
1950    shapes: &[&[SymDim]],
1951) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1952    const OP: &str = "tenferro-linalg.householder_qr_append_tangent";
1953    for shape in shapes {
1954        require_householder_rank2(OP, shape)?;
1955    }
1956    if dtypes.iter().any(|dtype| *dtype != dtypes[0]) {
1957        return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1958    }
1959    require_static_extent_equal(OP, "rows", &shapes[0][0], &shapes[1][0])?;
1960    require_static_extent_equal(OP, "left tangent", &shapes[0][0], &shapes[2][0])?;
1961    require_static_extent_equal(OP, "right tangent", &shapes[1][0], &shapes[3][0])?;
1962    require_static_extent_equal(OP, "anchor rows", &shapes[2][0], &shapes[3][0])?;
1963    Ok((
1964        dtypes[0],
1965        vec![
1966            shapes[2][0].clone(),
1967            shapes[2][1].clone() + shapes[3][1].clone(),
1968        ],
1969    ))
1970}
1971
1972fn householder_qr_split_tangent_meta(
1973    dtypes: &[DType],
1974    shapes: &[&[SymDim]],
1975    right: bool,
1976) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1977    const OP: &str = "tenferro-linalg.householder_qr_split_tangent";
1978    for shape in shapes {
1979        require_householder_rank2(OP, shape)?;
1980    }
1981    if dtypes.iter().any(|dtype| *dtype != dtypes[0]) {
1982        return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1983    }
1984    require_static_extent_equal(OP, "left rows", &shapes[0][0], &shapes[1][0])?;
1985    require_static_extent_equal(OP, "right rows", &shapes[0][0], &shapes[2][0])?;
1986    let expected_width = shapes[1][1].clone() + shapes[2][1].clone();
1987    require_static_extent_equal(OP, "width", &shapes[0][1], &expected_width)?;
1988    let selected = if right { shapes[2] } else { shapes[1] };
1989    Ok((dtypes[0], selected.to_vec()))
1990}
1991
1992fn eigh_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1993    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh", shape)?;
1994    Ok(vec![
1995        (singular_values_dtype(dtype), vector_shape(n.clone(), batch)),
1996        (dtype, matrix_shape(n.clone(), n, batch)),
1997    ])
1998}
1999
2000fn eigh_values_meta(
2001    dtype: DType,
2002    shape: &[SymDim],
2003) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
2004    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh_values", shape)?;
2005    Ok((singular_values_dtype(dtype), vector_shape(n, batch)))
2006}
2007
2008fn eig_meta(
2009    input_dtype: DType,
2010    shape: &[SymDim],
2011) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
2012    let dtype = eig_output_dtype(input_dtype);
2013    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig", shape)?;
2014    Ok(vec![
2015        (dtype, vector_shape(n.clone(), batch)),
2016        (dtype, matrix_shape(n.clone(), n, batch)),
2017    ])
2018}
2019
2020fn eig_values_meta(
2021    input_dtype: DType,
2022    shape: &[SymDim],
2023) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
2024    let dtype = eig_output_dtype(input_dtype);
2025    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig_values", shape)?;
2026    Ok((dtype, vector_shape(n, batch)))
2027}
2028
2029fn matrix_shape(rows: SymDim, cols: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
2030    let mut shape = vec![rows, cols];
2031    shape.extend_from_slice(batch);
2032    shape
2033}
2034
2035fn vector_shape(len: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
2036    let mut shape = vec![len];
2037    shape.extend_from_slice(batch);
2038    shape
2039}
2040
2041fn eig_output_dtype(dtype: DType) -> DType {
2042    match dtype {
2043        DType::F64 | DType::C64 => DType::C64,
2044        DType::F32 | DType::C32 => DType::C32,
2045        DType::I32 | DType::I64 | DType::Bool => DType::C64,
2046        // INVARIANT: linalg validates its input dtype before mapping outputs, so
2047        // an externally defined scalar never reaches this mapping.
2048        DType::External(_) => unreachable!("linalg validates its input dtype first"),
2049    }
2050}
2051
2052fn singular_values_dtype(dtype: DType) -> DType {
2053    match dtype {
2054        DType::C64 => DType::F64,
2055        DType::C32 => DType::F32,
2056        DType::External(id) => {
2057            // Singular values of an externally defined scalar have no declared
2058            // dtype, so the mapping keeps the identity rather than guessing one.
2059            DType::External(id)
2060        }
2061        other => other,
2062    }
2063}
2064
2065fn promote_dtypes(dtypes: &[DType]) -> DType {
2066    dtypes
2067        .iter()
2068        .copied()
2069        .reduce(tenferro_tensor::validate::promote_dtype)
2070        .unwrap_or(DType::F64)
2071}
2072
2073fn hash_dtype(hasher: &mut dyn Hasher, dtype: DType) {
2074    let tag = match dtype {
2075        DType::F64 => 0,
2076        DType::F32 => 1,
2077        DType::I64 => 2,
2078        DType::C64 => 3,
2079        DType::C32 => 4,
2080        DType::I32 => 5,
2081        DType::Bool => 6,
2082        DType::External(id) => {
2083            // The cache key must distinguish two externally defined scalars, so it
2084            // carries the scalar's own identity rather than one shared code.
2085            // `Hash` needs a sized hasher, so the identity is folded into a
2086            // concrete one first. It is stable within a process, which is the
2087            // scope of this cache.
2088            let mut identity = std::collections::hash_map::DefaultHasher::new();
2089            Hash::hash(&id, &mut identity);
2090            hasher.write_u8(7);
2091            hasher.write_u64(identity.finish());
2092            return;
2093        }
2094    };
2095    hasher.write_u8(tag);
2096}
2097
2098fn hash_svd_gauge(hasher: &mut dyn Hasher, gauge: SvdGauge) {
2099    let tag = match gauge {
2100        SvdGauge::Raw => 0,
2101        SvdGauge::CanonicalPivot => 1,
2102    };
2103    hasher.write_u8(tag);
2104}
2105
2106fn hash_svd_driver(hasher: &mut dyn Hasher, driver: SvdDriver) {
2107    let tag = match driver {
2108        SvdDriver::Auto => 0,
2109        SvdDriver::Gesvdj => 1,
2110        SvdDriver::Gesvd => 2,
2111        SvdDriver::Xgesvdp => 3,
2112    };
2113    hasher.write_u8(tag);
2114}
2115
2116fn hash_eigh_driver(hasher: &mut dyn Hasher, driver: EighDriver) {
2117    let tag = match driver {
2118        EighDriver::Auto => 0,
2119        EighDriver::Syevd => 1,
2120        EighDriver::Syevj => 2,
2121    };
2122    hasher.write_u8(tag);
2123}
2124
2125fn hash_eigh_gauge(hasher: &mut dyn Hasher, gauge: EighGauge) {
2126    let tag = match gauge {
2127        EighGauge::Raw => 0,
2128        EighGauge::CanonicalPivot => 1,
2129    };
2130    hasher.write_u8(tag);
2131}
2132
2133fn hash_qr_gauge(hasher: &mut dyn Hasher, gauge: QrGauge) {
2134    let tag = match gauge {
2135        QrGauge::Raw => 0,
2136        QrGauge::PositiveDiagonal => 1,
2137    };
2138    hasher.write_u8(tag);
2139}