Skip to main content

tenferro_linalg/
extension.rs

1use std::any::Any;
2use std::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::{ExtensionExecutionContext, ExtensionOp};
10use tenferro_tensor::{BackendSession, DType, Error, ErrorKind, Tensor, TensorBackend, TensorRead};
11
12#[cfg(feature = "cuda")]
13use tenferro_gpu::with_cuda_exec_session;
14
15use crate::backend::LinalgBackend;
16
17mod gauge;
18#[cfg(all(test, not(feature = "cuda")))]
19mod tests;
20
21pub(crate) use gauge::{apply_eigh_gauge, apply_qr_gauge};
22
23pub const LINALG_EXTENSION_FAMILY_ID: &str = "tenferro-linalg.linalg.v1";
24
25/// Default derivative regularization used by decomposition AD rules.
26///
27/// This epsilon is used only when differentiating decomposition formulas with
28/// repeated or nearly repeated spectral values. It is not a solver tolerance.
29///
30/// # Examples
31///
32/// ```rust
33/// use tenferro_linalg::{SvdOptions, DEFAULT_DECOMPOSITION_DERIVATIVE_EPS};
34///
35/// let options = SvdOptions::default();
36/// assert_eq!(options.derivative_eps, DEFAULT_DECOMPOSITION_DERIVATIVE_EPS);
37/// ```
38pub const DEFAULT_DECOMPOSITION_DERIVATIVE_EPS: f64 = 1e-12;
39
40/// Singular-vector gauge convention used by [`SvdOptions`].
41///
42/// # Examples
43///
44/// ```rust
45/// use tenferro_linalg::{SvdGauge, SvdOptions};
46///
47/// let options = SvdOptions::default().gauge(SvdGauge::CanonicalPivot);
48/// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
49/// ```
50#[derive(Clone, Copy, Debug, PartialEq, Eq)]
51pub enum SvdGauge {
52    /// Leave the backend's raw singular vector signs or phases unchanged.
53    Raw,
54    /// Make each left singular vector's max-absolute pivot entry positive-real
55    /// and adjust the matching `VT` row so reconstruction is preserved.
56    CanonicalPivot,
57}
58
59/// Eigenvector gauge convention used by [`EighOptions`].
60///
61/// # Examples
62///
63/// ```rust
64/// use tenferro_linalg::{EighGauge, EighOptions};
65///
66/// let options = EighOptions::default().gauge(EighGauge::CanonicalPivot);
67/// assert_eq!(options.gauge, EighGauge::CanonicalPivot);
68/// ```
69#[derive(Clone, Copy, Debug, PartialEq, Eq)]
70pub enum EighGauge {
71    /// Leave the backend's raw eigenvector signs or phases unchanged.
72    Raw,
73    /// Make each eigenvector's max-absolute pivot entry positive-real.
74    CanonicalPivot,
75}
76
77/// QR factor gauge convention used by [`QrOptions`].
78///
79/// # Examples
80///
81/// ```rust
82/// use tenferro_linalg::{QrGauge, QrOptions};
83///
84/// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
85/// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
86/// ```
87#[derive(Clone, Copy, Debug, PartialEq, Eq)]
88pub enum QrGauge {
89    /// Leave the backend's raw QR signs or phases unchanged.
90    Raw,
91    /// Make each `R` diagonal entry positive-real, compensating `Q`.
92    PositiveDiagonal,
93}
94
95/// Options for singular value decomposition.
96///
97/// # Examples
98///
99/// ```rust
100/// use tenferro_linalg::{SvdGauge, SvdOptions};
101///
102/// let options = SvdOptions::default()
103///     .gauge(SvdGauge::CanonicalPivot)
104///     .derivative_eps(1.0e-10);
105/// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
106/// assert_eq!(options.derivative_eps, 1.0e-10);
107/// ```
108#[derive(Clone, Copy, Debug, PartialEq)]
109pub struct SvdOptions {
110    /// Singular-vector gauge convention.
111    pub gauge: SvdGauge,
112    /// AD derivative regularization for repeated or nearly repeated singular values.
113    pub derivative_eps: f64,
114}
115
116impl Default for SvdOptions {
117    fn default() -> Self {
118        Self {
119            gauge: SvdGauge::Raw,
120            derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
121        }
122    }
123}
124
125impl SvdOptions {
126    /// Return options with the requested singular-vector gauge.
127    ///
128    /// # Examples
129    ///
130    /// ```rust
131    /// use tenferro_linalg::{SvdGauge, SvdOptions};
132    ///
133    /// let options = SvdOptions::default().gauge(SvdGauge::CanonicalPivot);
134    /// assert_eq!(options.gauge, SvdGauge::CanonicalPivot);
135    /// ```
136    pub fn gauge(mut self, gauge: SvdGauge) -> Self {
137        self.gauge = gauge;
138        self
139    }
140
141    /// Return options with an explicit derivative epsilon.
142    ///
143    /// # Examples
144    ///
145    /// ```rust
146    /// use tenferro_linalg::SvdOptions;
147    ///
148    /// let options = SvdOptions::default().derivative_eps(1.0e-9);
149    /// assert_eq!(options.derivative_eps, 1.0e-9);
150    /// ```
151    pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
152        self.derivative_eps = derivative_eps;
153        self
154    }
155}
156
157/// Options for Hermitian eigenvalue decomposition.
158///
159/// # Examples
160///
161/// ```rust
162/// use tenferro_linalg::EighOptions;
163///
164/// let options = EighOptions::default().derivative_eps(1.0e-10);
165/// assert_eq!(options.derivative_eps, 1.0e-10);
166/// ```
167#[derive(Clone, Copy, Debug, PartialEq)]
168pub struct EighOptions {
169    /// Eigenvector gauge convention.
170    pub gauge: EighGauge,
171    /// AD derivative regularization for repeated or nearly repeated eigenvalues.
172    pub derivative_eps: f64,
173}
174
175impl Default for EighOptions {
176    fn default() -> Self {
177        Self {
178            gauge: EighGauge::Raw,
179            derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
180        }
181    }
182}
183
184impl EighOptions {
185    /// Return options with the requested eigenvector gauge.
186    ///
187    /// # Examples
188    ///
189    /// ```rust
190    /// use tenferro_linalg::{EighGauge, EighOptions};
191    ///
192    /// let options = EighOptions::default().gauge(EighGauge::CanonicalPivot);
193    /// assert_eq!(options.gauge, EighGauge::CanonicalPivot);
194    /// ```
195    pub fn gauge(mut self, gauge: EighGauge) -> Self {
196        self.gauge = gauge;
197        self
198    }
199
200    /// Return options with an explicit derivative epsilon.
201    ///
202    /// # Examples
203    ///
204    /// ```rust
205    /// use tenferro_linalg::EighOptions;
206    ///
207    /// let options = EighOptions::default().derivative_eps(1.0e-9);
208    /// assert_eq!(options.derivative_eps, 1.0e-9);
209    /// ```
210    pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
211        self.derivative_eps = derivative_eps;
212        self
213    }
214}
215
216/// Options for QR decomposition.
217///
218/// # Examples
219///
220/// ```rust
221/// use tenferro_linalg::{QrGauge, QrOptions};
222///
223/// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
224/// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
225/// ```
226#[derive(Clone, Copy, Debug, PartialEq, Eq)]
227pub struct QrOptions {
228    /// QR sign or phase convention.
229    pub gauge: QrGauge,
230}
231
232impl Default for QrOptions {
233    fn default() -> Self {
234        Self {
235            gauge: QrGauge::Raw,
236        }
237    }
238}
239
240impl QrOptions {
241    /// Return options with the requested QR gauge.
242    ///
243    /// # Examples
244    ///
245    /// ```rust
246    /// use tenferro_linalg::{QrGauge, QrOptions};
247    ///
248    /// let options = QrOptions::default().gauge(QrGauge::PositiveDiagonal);
249    /// assert_eq!(options.gauge, QrGauge::PositiveDiagonal);
250    /// ```
251    pub fn gauge(mut self, gauge: QrGauge) -> Self {
252        self.gauge = gauge;
253        self
254    }
255}
256
257pub(crate) fn validate_derivative_eps(
258    op: &'static str,
259    derivative_eps: f64,
260) -> tenferro_tensor::Result<()> {
261    if derivative_eps.is_finite() && derivative_eps > 0.0 {
262        Ok(())
263    } else {
264        Err(Error::invalid_argument(
265            op,
266            "derivative_eps",
267            format!("must be positive and finite, got {derivative_eps}"),
268        ))
269    }
270}
271
272#[derive(Clone, Copy, Debug, PartialEq)]
273#[doc(hidden)]
274pub(crate) enum LinalgOp {
275    Cholesky,
276    Lu,
277    LuFactor,
278    LuSolvePrepared {
279        transpose_a: bool,
280        conjugate_a: bool,
281    },
282    SignDetFromLuFactor,
283    LogAbsDetFromLuFactor,
284    FullPivLu,
285    FullPivLuSolve {
286        transpose_a: bool,
287    },
288    Svd {
289        derivative_eps: f64,
290        gauge: SvdGauge,
291    },
292    /// Full-matrices SVD: `U` is `m x m` and `Vh` is `n x n`, so the trailing
293    /// `Vh` rows span the input's right nullspace. Value-only: AD is
294    /// intentionally unsupported (see the linalg AD support manifest).
295    SvdFull,
296    SvdVals {
297        derivative_eps: f64,
298    },
299    Qr {
300        gauge: QrGauge,
301    },
302    Eigh {
303        derivative_eps: f64,
304        gauge: EighGauge,
305    },
306    EighVals {
307        derivative_eps: f64,
308    },
309    Eig {
310        input_dtype: DType,
311    },
312    EigVals {
313        input_dtype: DType,
314    },
315    TriangularSolve {
316        left_side: bool,
317        lower: bool,
318        transpose_a: bool,
319        unit_diagonal: bool,
320    },
321}
322
323impl LinalgOp {
324    fn output_count(self) -> usize {
325        match self {
326            Self::Cholesky
327            | Self::EighVals { .. }
328            | Self::EigVals { .. }
329            | Self::FullPivLuSolve { .. }
330            | Self::LogAbsDetFromLuFactor
331            | Self::LuSolvePrepared { .. }
332            | Self::SignDetFromLuFactor
333            | Self::SvdVals { .. }
334            | Self::TriangularSolve { .. } => 1,
335            Self::Svd { .. } | Self::SvdFull => 3,
336            Self::Qr { .. } | Self::Eigh { .. } | Self::Eig { .. } => 2,
337            Self::LuFactor => 3,
338            Self::Lu => 4,
339            Self::FullPivLu => 5,
340        }
341    }
342
343    fn input_count(self) -> usize {
344        match self {
345            Self::FullPivLuSolve { .. } | Self::TriangularSolve { .. } => 2,
346            Self::LogAbsDetFromLuFactor => 2,
347            Self::SignDetFromLuFactor => 3,
348            Self::LuSolvePrepared { .. } => 4,
349            _ => 1,
350        }
351    }
352
353    fn tag(self) -> u8 {
354        match self {
355            Self::Cholesky => 0,
356            Self::Lu => 1,
357            Self::FullPivLu => 2,
358            Self::FullPivLuSolve { .. } => 3,
359            Self::Svd { .. } => 4,
360            Self::Qr { .. } => 5,
361            Self::Eigh { .. } => 6,
362            Self::Eig { .. } => 7,
363            Self::TriangularSolve { .. } => 9,
364            Self::LuFactor => 10,
365            Self::LuSolvePrepared { .. } => 11,
366            Self::SvdVals { .. } => 12,
367            Self::EighVals { .. } => 13,
368            Self::EigVals { .. } => 14,
369            Self::SvdFull => 15,
370            Self::LogAbsDetFromLuFactor => 16,
371            Self::SignDetFromLuFactor => 17,
372        }
373    }
374}
375
376#[derive(Clone, Debug, PartialEq)]
377#[doc(hidden)]
378pub(crate) struct LinalgExtensionOp {
379    op: LinalgOp,
380}
381
382impl LinalgExtensionOp {
383    pub(crate) fn new(op: LinalgOp) -> Self {
384        Self { op }
385    }
386
387    pub(crate) fn op(&self) -> LinalgOp {
388        self.op
389    }
390}
391
392impl ExtensionOp for LinalgExtensionOp {
393    fn family_id(&self) -> &'static str {
394        LINALG_EXTENSION_FAMILY_ID
395    }
396
397    fn payload_hash(&self, hasher: &mut dyn Hasher) {
398        hasher.write_u8(self.op.tag());
399        match self.op {
400            LinalgOp::Svd {
401                derivative_eps,
402                gauge,
403            } => {
404                hasher.write_u64(derivative_eps.to_bits());
405                hash_svd_gauge(hasher, gauge);
406            }
407            LinalgOp::SvdVals { derivative_eps } | LinalgOp::EighVals { derivative_eps } => {
408                hasher.write_u64(derivative_eps.to_bits());
409            }
410            LinalgOp::Qr { gauge } => {
411                hash_qr_gauge(hasher, gauge);
412            }
413            LinalgOp::Eigh {
414                derivative_eps,
415                gauge,
416            } => {
417                hasher.write_u64(derivative_eps.to_bits());
418                hash_eigh_gauge(hasher, gauge);
419            }
420            LinalgOp::Eig { input_dtype } | LinalgOp::EigVals { input_dtype } => {
421                hash_dtype(hasher, input_dtype);
422            }
423            LinalgOp::FullPivLuSolve { transpose_a } => {
424                hasher.write_u8(u8::from(transpose_a));
425            }
426            LinalgOp::LuSolvePrepared {
427                transpose_a,
428                conjugate_a,
429            } => {
430                hasher.write_u8(u8::from(transpose_a));
431                hasher.write_u8(u8::from(conjugate_a));
432            }
433            LinalgOp::TriangularSolve {
434                left_side,
435                lower,
436                transpose_a,
437                unit_diagonal,
438            } => {
439                hasher.write_u8(u8::from(left_side));
440                hasher.write_u8(u8::from(lower));
441                hasher.write_u8(u8::from(transpose_a));
442                hasher.write_u8(u8::from(unit_diagonal));
443            }
444            LinalgOp::Cholesky
445            | LinalgOp::Lu
446            | LinalgOp::LuFactor
447            | LinalgOp::LogAbsDetFromLuFactor
448            | LinalgOp::SignDetFromLuFactor
449            | LinalgOp::FullPivLu
450            | LinalgOp::SvdFull => {}
451        }
452    }
453
454    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
455        other
456            .as_any()
457            .downcast_ref::<Self>()
458            .is_some_and(|that| self == that)
459    }
460
461    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
462        Arc::new(self.clone())
463    }
464
465    fn as_any(&self) -> &dyn Any {
466        self
467    }
468
469    fn input_count(&self) -> usize {
470        self.op.input_count()
471    }
472
473    fn output_count(&self) -> usize {
474        self.op.output_count()
475    }
476
477    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
478        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
479    }
480
481    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
482        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
483    }
484
485    fn prune_outputs(&self, live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
486        match self.op {
487            LinalgOp::Svd { derivative_eps, .. } if live_outputs == [false, true, false] => {
488                Some(Arc::new(Self::new(LinalgOp::SvdVals { derivative_eps })))
489            }
490            LinalgOp::Eigh { derivative_eps, .. } if live_outputs == [true, false] => {
491                Some(Arc::new(Self::new(LinalgOp::EighVals { derivative_eps })))
492            }
493            LinalgOp::Eig { input_dtype } if live_outputs == [true, false] => {
494                Some(Arc::new(Self::new(LinalgOp::EigVals { input_dtype })))
495            }
496            _ => None,
497        }
498    }
499
500    fn infer_output_meta(
501        &self,
502        ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
503    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
504        let input_dtypes = (0..self.input_count())
505            .map(|input| ctx.input_dtype(input))
506            .collect::<Result<Vec<_>, _>>()?;
507        let input_shapes = (0..self.input_count())
508            .map(|input| ctx.input_shape(input))
509            .collect::<Result<Vec<_>, _>>()?;
510        let metas = match self.op {
511            LinalgOp::Cholesky => {
512                require_matrix_meta("tenferro-linalg.cholesky", input_shapes[0])?;
513                vec![(promote_dtypes(&input_dtypes), input_shapes[0].to_vec())]
514            }
515            LinalgOp::FullPivLuSolve { .. } => {
516                require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[0])?;
517                require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[1])?;
518                vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
519            }
520            LinalgOp::TriangularSolve { .. } => {
521                require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[0])?;
522                require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[1])?;
523                vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
524            }
525            LinalgOp::LuSolvePrepared { .. } => {
526                require_matrix_meta("tenferro-linalg.lu_solve_prepared_lu", input_shapes[0])?;
527                require_matrix_meta("tenferro-linalg.lu_solve_prepared_rhs", input_shapes[3])?;
528                vec![(
529                    promote_dtypes(&[input_dtypes[0], input_dtypes[3]]),
530                    input_shapes[3].to_vec(),
531                )]
532            }
533            LinalgOp::Lu => lu_meta(input_dtypes[0], input_shapes[0])?,
534            LinalgOp::LuFactor => lu_factor_meta(input_dtypes[0], input_shapes[0])?,
535            LinalgOp::SignDetFromLuFactor => {
536                vec![signdet_from_lu_factor_meta(
537                    input_dtypes[0],
538                    input_shapes[0],
539                    input_shapes[1],
540                    input_shapes[2],
541                )?]
542            }
543            LinalgOp::LogAbsDetFromLuFactor => {
544                vec![logabsdet_from_lu_factor_meta(
545                    input_dtypes[0],
546                    input_shapes[0],
547                    input_shapes[1],
548                )?]
549            }
550            LinalgOp::FullPivLu => full_piv_lu_meta(input_dtypes[0], input_shapes[0])?,
551            LinalgOp::Svd { .. } => svd_meta(input_dtypes[0], input_shapes[0])?,
552            LinalgOp::SvdFull => svd_full_meta(input_dtypes[0], input_shapes[0])?,
553            LinalgOp::SvdVals { .. } => {
554                vec![svd_values_meta(input_dtypes[0], input_shapes[0])?]
555            }
556            LinalgOp::Qr { .. } => qr_meta(input_dtypes[0], input_shapes[0])?,
557            LinalgOp::Eigh { .. } => eigh_meta(input_dtypes[0], input_shapes[0])?,
558            LinalgOp::EighVals { .. } => vec![eigh_values_meta(input_dtypes[0], input_shapes[0])?],
559            LinalgOp::Eig { input_dtype } => eig_meta(input_dtype, input_shapes[0])?,
560            LinalgOp::EigVals { input_dtype } => {
561                vec![eig_values_meta(input_dtype, input_shapes[0])?]
562            }
563        };
564        Ok(metas)
565    }
566}
567
568pub(crate) fn execute_linalg_extension_reads<B: BackendSession + ?Sized>(
569    op: &LinalgExtensionOp,
570    inputs: &[TensorRead<'_>],
571    ctx: &mut ExtensionExecutionContext<'_, B>,
572) -> tenferro_tensor::Result<Vec<Tensor>> {
573    execute_linalg_extension_reads_on_session(op, inputs, ctx.backend_mut())
574}
575
576pub(crate) fn execute_linalg_extension_reads_owner<B: TensorBackend>(
577    op: &LinalgExtensionOp,
578    inputs: &[TensorRead<'_>],
579    ctx: &mut ExtensionExecutionContext<'_, B>,
580) -> tenferro_tensor::Result<Vec<Tensor>> {
581    let (backend, caches) = ctx.parts_mut();
582    backend.with_backend_session(|session| {
583        let mut session_ctx = ExtensionExecutionContext::new(session, caches);
584        execute_linalg_extension_reads(op, inputs, &mut session_ctx)
585    })
586}
587
588fn execute_linalg_extension_reads_on_session<B: BackendSession + ?Sized>(
589    op: &LinalgExtensionOp,
590    inputs: &[TensorRead<'_>],
591    session: &mut B,
592) -> tenferro_tensor::Result<Vec<Tensor>> {
593    if let Some(result) = with_cpu_exec_session(session, |session| {
594        execute_linalg_extension_reads_in_session(op, inputs, session)
595    }) {
596        return result;
597    }
598    #[cfg(feature = "cuda")]
599    if let Some(result) = with_cuda_exec_session(session, |session| {
600        execute_linalg_extension_reads_in_session(op, inputs, session)
601    }) {
602        return result;
603    }
604    Err(Error::unsupported(
605        "linalg_extension",
606        "selected backend session does not expose a linalg execution capability",
607    ))
608}
609
610fn execute_linalg_extension_reads_in_session<S: LinalgBackend>(
611    op: &LinalgExtensionOp,
612    inputs: &[TensorRead<'_>],
613    session: &mut S,
614) -> tenferro_tensor::Result<Vec<Tensor>> {
615    if op.op() == LinalgOp::Cholesky {
616        return Ok(vec![session.cholesky_read(inputs[0].clone())?]);
617    }
618    if let LinalgOp::TriangularSolve {
619        left_side,
620        lower,
621        transpose_a,
622        unit_diagonal,
623    } = op.op()
624    {
625        match session.triangular_solve_read(
626            inputs[0].clone(),
627            inputs[1].clone(),
628            left_side,
629            lower,
630            transpose_a,
631            unit_diagonal,
632        ) {
633            Ok(output) => return Ok(vec![output]),
634            Err(error) if error.kind() == ErrorKind::Unsupported => {}
635            Err(error) => return Err(error),
636        }
637    }
638
639    // Linalg kernels operate on compact tensors; materialization is explicit
640    // here so borrowed views cannot bypass provider errors.
641    let materialized_inputs = inputs
642        .iter()
643        .cloned()
644        .map(|input| session.to_contiguous_read(input))
645        .collect::<tenferro_tensor::Result<Vec<_>>>()?;
646    let input_refs: Vec<&Tensor> = materialized_inputs.iter().collect();
647    execute_linalg(op.op(), &input_refs, session)
648}
649
650fn linalg_session_supported<B: BackendSession + 'static>(op: &LinalgExtensionOp) -> bool {
651    matches!(op.op(), LinalgOp::LuSolvePrepared { .. })
652        && std::any::TypeId::of::<B>() == std::any::TypeId::of::<tenferro_cpu::CpuBackend>()
653}
654
655fn execute_linalg_extension_in_session(
656    op: &LinalgExtensionOp,
657    session: &mut dyn BackendSession,
658    _extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
659    inputs: &[TensorRead<'_>],
660) -> tenferro_tensor::Result<Vec<Tensor>> {
661    if !matches!(op.op(), LinalgOp::LuSolvePrepared { .. }) {
662        return Err(Error::unsupported(
663            "linalg_extension",
664            "linalg operation has no scheduler-session implementation",
665        ));
666    };
667    if inputs.len() != 4 {
668        return Err(Error::invalid_argument(
669            "linalg_extension",
670            "inputs",
671            format!(
672                "LuSolvePrepared session execution expected 4 inputs, got {}",
673                inputs.len()
674            ),
675        ));
676    }
677    execute_linalg_extension_reads_on_session(op, inputs, session)
678}
679
680define_extension_runtime! {
681    runtime = LinalgRuntime,
682    family_id = LINALG_EXTENSION_FAMILY_ID,
683    op_type = LinalgExtensionOp,
684    execute = execute_linalg_extension_reads_owner,
685    execute_reads = execute_linalg_extension_reads_owner,
686    execute_in_session = execute_linalg_extension_in_session,
687    session_supported = linalg_session_supported,
688    backend_bound = TensorBackend,
689}
690
691fn execute_linalg<B: LinalgBackend>(
692    op: LinalgOp,
693    inputs: &[&Tensor],
694    backend: &mut B,
695) -> tenferro_tensor::Result<Vec<Tensor>> {
696    match op {
697        LinalgOp::Cholesky => Ok(vec![backend.cholesky(inputs[0])?]),
698        LinalgOp::Lu => backend.lu(inputs[0]),
699        LinalgOp::LuFactor => backend.lu_factor(inputs[0]),
700        LinalgOp::SignDetFromLuFactor => Ok(vec![signdet_from_lu_factor(
701            inputs[0].dtype(),
702            inputs[1],
703            inputs[2],
704            backend,
705        )?]),
706        LinalgOp::LogAbsDetFromLuFactor => Ok(vec![logabsdet_from_lu_factor(inputs[1], backend)?]),
707        LinalgOp::LuSolvePrepared {
708            transpose_a,
709            conjugate_a,
710        } => Ok(vec![backend.lu_solve_prepared(
711            inputs[0],
712            inputs[1],
713            inputs[2],
714            inputs[3],
715            transpose_a,
716            conjugate_a,
717        )?]),
718        LinalgOp::FullPivLu => backend.full_piv_lu(inputs[0]),
719        LinalgOp::FullPivLuSolve { transpose_a } => Ok(vec![backend.full_piv_lu_solve(
720            inputs[0],
721            inputs[1],
722            transpose_a,
723        )?]),
724        LinalgOp::Svd {
725            derivative_eps,
726            gauge,
727        } => backend.svd_with_options(
728            inputs[0],
729            SvdOptions {
730                derivative_eps,
731                gauge,
732            },
733        ),
734        LinalgOp::SvdFull => backend.svd_full(inputs[0]),
735        LinalgOp::SvdVals { .. } => Ok(vec![backend.svd_values(inputs[0])?]),
736        LinalgOp::Qr { gauge } => backend.qr_with_options(inputs[0], QrOptions { gauge }),
737        LinalgOp::Eigh {
738            derivative_eps,
739            gauge,
740        } => backend.eigh_with_options(
741            inputs[0],
742            EighOptions {
743                derivative_eps,
744                gauge,
745            },
746        ),
747        LinalgOp::EighVals { .. } => Ok(vec![backend.eigh_values(inputs[0])?]),
748        LinalgOp::Eig { .. } => backend.eig(inputs[0]),
749        LinalgOp::EigVals { .. } => Ok(vec![backend.eig_values(inputs[0])?]),
750        LinalgOp::TriangularSolve {
751            left_side,
752            lower,
753            transpose_a,
754            unit_diagonal,
755        } => Ok(vec![backend.triangular_solve(
756            inputs[0],
757            inputs[1],
758            left_side,
759            lower,
760            transpose_a,
761            unit_diagonal,
762        )?]),
763    }
764}
765
766fn signdet_from_lu_factor<B: LinalgBackend + ?Sized>(
767    input_dtype: DType,
768    packed_lu: &Tensor,
769    parity: &Tensor,
770    backend: &mut B,
771) -> tenferro_tensor::Result<Tensor> {
772    let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
773    let det_u = backend.reduce_prod_read(TensorRead::from_tensor(&diag), &[0])?;
774    let det = backend.mul_read(
775        TensorRead::from_tensor(parity),
776        TensorRead::from_tensor(&det_u),
777    )?;
778    if matches!(input_dtype, DType::C32 | DType::C64) {
779        let abs = backend.abs_read(TensorRead::from_tensor(&det))?;
780        let abs = backend.convert(&abs, input_dtype)?;
781        backend.div_read(TensorRead::from_tensor(&det), TensorRead::from_tensor(&abs))
782    } else {
783        backend.sign_read(TensorRead::from_tensor(&det))
784    }
785}
786
787fn logabsdet_from_lu_factor<B: LinalgBackend + ?Sized>(
788    packed_lu: &Tensor,
789    backend: &mut B,
790) -> tenferro_tensor::Result<Tensor> {
791    let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
792    let abs = backend.abs_read(TensorRead::from_tensor(&diag))?;
793    let log = backend.log_read(TensorRead::from_tensor(&abs))?;
794    backend.reduce_sum_read(TensorRead::from_tensor(&log), &[0])
795}
796
797pub(crate) fn apply_svd_gauge(
798    gauge: SvdGauge,
799    outputs: &mut [Tensor],
800) -> tenferro_tensor::Result<()> {
801    match gauge {
802        SvdGauge::Raw => Ok(()),
803        SvdGauge::CanonicalPivot => apply_canonical_pivot_svd_gauge(outputs),
804    }
805}
806
807fn apply_canonical_pivot_svd_gauge(outputs: &mut [Tensor]) -> tenferro_tensor::Result<()> {
808    if outputs.len() != 3 {
809        return Err(Error::invalid_argument(
810            "tenferro-linalg.svd",
811            "outputs",
812            format!(
813                "canonical SVD gauge expected three outputs, got {}",
814                outputs.len()
815            ),
816        ));
817    }
818
819    let (u_slice, rest) = outputs.split_at_mut(1);
820    let (singular_slice, vt_slice) = rest.split_at_mut(1);
821    let u = &mut u_slice[0];
822    let singular_values = &singular_slice[0];
823    let vt = &mut vt_slice[0];
824    let u_shape = u.shape().to_vec();
825    let s_shape = singular_values.shape().to_vec();
826    let vt_shape = vt.shape().to_vec();
827    if u_shape.len() < 2 || vt_shape.len() < 2 || s_shape.is_empty() {
828        return Err(Error::invalid_argument(
829            "tenferro-linalg.svd",
830            "outputs",
831            format!(
832                "canonical SVD gauge expected U rank >= 2, S rank >= 1, VT rank >= 2; got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
833            ),
834        ));
835    }
836
837    let m = u_shape[0];
838    let k = u_shape[1];
839    let n = vt_shape[1];
840    if s_shape[0] != k
841        || vt_shape[0] != k
842        || u_shape[2..] != vt_shape[2..]
843        || s_shape[1..] != u_shape[2..]
844    {
845        return Err(Error::invalid_argument(
846            "tenferro-linalg.svd",
847            "outputs",
848            format!(
849                "canonical SVD gauge expected compatible compact SVD shapes, got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
850            ),
851        ));
852    }
853    let layout = canonical_svd_gauge_layout(m, k, n, &u_shape[2..])?;
854
855    match (u, vt) {
856        (Tensor::F64(u), Tensor::F64(vt)) => {
857            canonicalize_svd_gauge_f64(u.host_data_mut()?, vt.host_data_mut()?, layout)
858        }
859        (Tensor::F32(u), Tensor::F32(vt)) => {
860            canonicalize_svd_gauge_f32(u.host_data_mut()?, vt.host_data_mut()?, layout)
861        }
862        (Tensor::C64(u), Tensor::C64(vt)) => {
863            canonicalize_svd_gauge_c64(u.host_data_mut()?, vt.host_data_mut()?, layout)
864        }
865        (Tensor::C32(u), Tensor::C32(vt)) => {
866            canonicalize_svd_gauge_c32(u.host_data_mut()?, vt.host_data_mut()?, layout)
867        }
868        (u, vt) => Err(Error::dtype_mismatch(
869            "tenferro-linalg.svd",
870            u.dtype(),
871            vt.dtype(),
872        )),
873    }
874}
875
876#[derive(Clone, Copy, Debug, PartialEq, Eq)]
877struct CanonicalSvdGaugeLayout {
878    m: usize,
879    k: usize,
880    batch_count: usize,
881    u_batch_len: usize,
882    vt_batch_len: usize,
883    u_len: usize,
884    vt_len: usize,
885}
886
887impl CanonicalSvdGaugeLayout {
888    fn validate_storage(self, u_len: usize, vt_len: usize) -> tenferro_tensor::Result<()> {
889        if u_len != self.u_len {
890            return Err(Error::invalid_argument(
891                "tenferro-linalg.svd",
892                "U storage",
893                format!(
894                    "canonical SVD gauge expected U storage length {}, got {u_len}",
895                    self.u_len
896                ),
897            ));
898        }
899        if vt_len != self.vt_len {
900            return Err(Error::invalid_argument(
901                "tenferro-linalg.svd",
902                "VT storage",
903                format!(
904                    "canonical SVD gauge expected VT storage length {}, got {vt_len}",
905                    self.vt_len
906                ),
907            ));
908        }
909        Ok(())
910    }
911}
912
913fn canonical_svd_gauge_layout(
914    m: usize,
915    k: usize,
916    n: usize,
917    batch_shape: &[usize],
918) -> tenferro_tensor::Result<CanonicalSvdGaugeLayout> {
919    let batch_count = tenferro_tensor::validate::checked_shape_product(
920        "tenferro-linalg.svd",
921        "canonical SVD batch",
922        batch_shape,
923    )?;
924    let u_batch_len = tenferro_tensor::validate::checked_shape_product(
925        "tenferro-linalg.svd",
926        "canonical SVD U batch",
927        &[m, k],
928    )?;
929    let vt_batch_len = tenferro_tensor::validate::checked_shape_product(
930        "tenferro-linalg.svd",
931        "canonical SVD VT batch",
932        &[k, n],
933    )?;
934    let u_len = tenferro_tensor::validate::checked_shape_product(
935        "tenferro-linalg.svd",
936        "canonical SVD U storage",
937        &[u_batch_len, batch_count],
938    )?;
939    let vt_len = tenferro_tensor::validate::checked_shape_product(
940        "tenferro-linalg.svd",
941        "canonical SVD VT storage",
942        &[vt_batch_len, batch_count],
943    )?;
944    Ok(CanonicalSvdGaugeLayout {
945        m,
946        k,
947        batch_count,
948        u_batch_len,
949        vt_batch_len,
950        u_len,
951        vt_len,
952    })
953}
954
955fn canonicalize_svd_gauge_f64(
956    u: &mut [f64],
957    vt: &mut [f64],
958    layout: CanonicalSvdGaugeLayout,
959) -> tenferro_tensor::Result<()> {
960    layout.validate_storage(u.len(), vt.len())?;
961    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
962        return Ok(());
963    }
964    for (u_batch, vt_batch) in u
965        .chunks_exact_mut(layout.u_batch_len)
966        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
967    {
968        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
969            let pivot = max_abs_pivot_f64(u_column);
970            let pivot_value = u_column[pivot];
971            if pivot_value < 0.0 {
972                for value in u_column {
973                    *value = -*value;
974                }
975                for vt_column in vt_batch.chunks_exact_mut(layout.k) {
976                    vt_column[col] = -vt_column[col];
977                }
978            }
979        }
980    }
981    Ok(())
982}
983
984fn canonicalize_svd_gauge_f32(
985    u: &mut [f32],
986    vt: &mut [f32],
987    layout: CanonicalSvdGaugeLayout,
988) -> tenferro_tensor::Result<()> {
989    layout.validate_storage(u.len(), vt.len())?;
990    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
991        return Ok(());
992    }
993    for (u_batch, vt_batch) in u
994        .chunks_exact_mut(layout.u_batch_len)
995        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
996    {
997        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
998            let pivot = max_abs_pivot_f32(u_column);
999            let pivot_value = u_column[pivot];
1000            if pivot_value < 0.0 {
1001                for value in u_column {
1002                    *value = -*value;
1003                }
1004                for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1005                    vt_column[col] = -vt_column[col];
1006                }
1007            }
1008        }
1009    }
1010    Ok(())
1011}
1012
1013fn canonicalize_svd_gauge_c64(
1014    u: &mut [Complex64],
1015    vt: &mut [Complex64],
1016    layout: CanonicalSvdGaugeLayout,
1017) -> tenferro_tensor::Result<()> {
1018    layout.validate_storage(u.len(), vt.len())?;
1019    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1020        return Ok(());
1021    }
1022    for (u_batch, vt_batch) in u
1023        .chunks_exact_mut(layout.u_batch_len)
1024        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1025    {
1026        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1027            let pivot = max_abs_pivot_c64(u_column);
1028            let pivot_value = u_column[pivot];
1029            let pivot_norm = pivot_value.norm();
1030            if pivot_norm == 0.0 {
1031                continue;
1032            }
1033            let phase = pivot_value.conj() / pivot_norm;
1034            let vt_phase = phase.conj();
1035            for value in u_column {
1036                *value *= phase;
1037            }
1038            for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1039                vt_column[col] *= vt_phase;
1040            }
1041        }
1042    }
1043    Ok(())
1044}
1045
1046fn canonicalize_svd_gauge_c32(
1047    u: &mut [Complex32],
1048    vt: &mut [Complex32],
1049    layout: CanonicalSvdGaugeLayout,
1050) -> tenferro_tensor::Result<()> {
1051    layout.validate_storage(u.len(), vt.len())?;
1052    if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1053        return Ok(());
1054    }
1055    for (u_batch, vt_batch) in u
1056        .chunks_exact_mut(layout.u_batch_len)
1057        .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1058    {
1059        for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1060            let pivot = max_abs_pivot_c32(u_column);
1061            let pivot_value = u_column[pivot];
1062            let pivot_norm = pivot_value.norm();
1063            if pivot_norm == 0.0 {
1064                continue;
1065            }
1066            let phase = pivot_value.conj() / pivot_norm;
1067            let vt_phase = phase.conj();
1068            for value in u_column {
1069                *value *= phase;
1070            }
1071            for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1072                vt_column[col] *= vt_phase;
1073            }
1074        }
1075    }
1076    Ok(())
1077}
1078
1079fn max_abs_pivot_f64(u_column: &[f64]) -> usize {
1080    let mut pivot = 0;
1081    let mut pivot_abs = u_column[0].abs();
1082    for (row, value) in u_column.iter().enumerate().skip(1) {
1083        let candidate_abs = value.abs();
1084        if candidate_abs > pivot_abs {
1085            pivot = row;
1086            pivot_abs = candidate_abs;
1087        }
1088    }
1089    pivot
1090}
1091
1092fn max_abs_pivot_f32(u_column: &[f32]) -> usize {
1093    let mut pivot = 0;
1094    let mut pivot_abs = u_column[0].abs();
1095    for (row, value) in u_column.iter().enumerate().skip(1) {
1096        let candidate_abs = value.abs();
1097        if candidate_abs > pivot_abs {
1098            pivot = row;
1099            pivot_abs = candidate_abs;
1100        }
1101    }
1102    pivot
1103}
1104
1105fn max_abs_pivot_c64(u_column: &[Complex64]) -> usize {
1106    let mut pivot = 0;
1107    let mut pivot_abs = u_column[0].norm_sqr();
1108    for (row, value) in u_column.iter().enumerate().skip(1) {
1109        let candidate_abs = value.norm_sqr();
1110        if candidate_abs > pivot_abs {
1111            pivot = row;
1112            pivot_abs = candidate_abs;
1113        }
1114    }
1115    pivot
1116}
1117
1118fn max_abs_pivot_c32(u_column: &[Complex32]) -> usize {
1119    let mut pivot = 0;
1120    let mut pivot_abs = u_column[0].norm_sqr();
1121    for (row, value) in u_column.iter().enumerate().skip(1) {
1122        let candidate_abs = value.norm_sqr();
1123        if candidate_abs > pivot_abs {
1124            pivot = row;
1125            pivot_abs = candidate_abs;
1126        }
1127    }
1128    pivot
1129}
1130
1131fn require_matrix_meta(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1132    if shape.len() < 2 {
1133        return Err(Error::rank_mismatch(op, 2, shape.len()));
1134    }
1135    Ok(())
1136}
1137
1138fn matrix_meta_parts<'a>(
1139    op: &'static str,
1140    shape: &'a [SymDim],
1141) -> tenferro_tensor::Result<(SymDim, SymDim, &'a [SymDim])> {
1142    require_matrix_meta(op, shape)?;
1143    Ok((shape[0].clone(), shape[1].clone(), &shape[2..]))
1144}
1145
1146fn lu_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1147    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu", shape)?;
1148    let k = m.clone().min(n.clone());
1149    Ok(vec![
1150        (dtype, matrix_shape(m.clone(), m, batch)),
1151        (dtype, matrix_shape(shape[0].clone(), k.clone(), batch)),
1152        (dtype, matrix_shape(k, n, batch)),
1153        (dtype, batch.to_vec()),
1154    ])
1155}
1156
1157fn lu_factor_meta(
1158    dtype: DType,
1159    shape: &[SymDim],
1160) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1161    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu_factor", shape)?;
1162    let k = m.min(n);
1163    Ok(vec![
1164        (dtype, shape.to_vec()),
1165        (DType::I32, vector_shape(k, batch)),
1166        (dtype, batch.to_vec()),
1167    ])
1168}
1169
1170fn signdet_from_lu_factor_meta(
1171    input_dtype: DType,
1172    input_shape: &[SymDim],
1173    packed_shape: &[SymDim],
1174    parity_shape: &[SymDim],
1175) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1176    let (_, _, batch) = matrix_meta_parts("tenferro-linalg.signdet_from_lu_factor", input_shape)?;
1177    require_matrix_meta(
1178        "tenferro-linalg.signdet_from_lu_factor_packed",
1179        packed_shape,
1180    )?;
1181    if parity_shape.len() != batch.len() {
1182        return Err(Error::rank_mismatch(
1183            "tenferro-linalg.signdet_from_lu_factor_parity",
1184            batch.len(),
1185            parity_shape.len(),
1186        ));
1187    }
1188    Ok((input_dtype, batch.to_vec()))
1189}
1190
1191fn logabsdet_from_lu_factor_meta(
1192    input_dtype: DType,
1193    input_shape: &[SymDim],
1194    packed_shape: &[SymDim],
1195) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1196    let (_, _, batch) = matrix_meta_parts("tenferro-linalg.logabsdet_from_lu_factor", input_shape)?;
1197    require_matrix_meta(
1198        "tenferro-linalg.logabsdet_from_lu_factor_packed",
1199        packed_shape,
1200    )?;
1201    Ok((singular_values_dtype(input_dtype), batch.to_vec()))
1202}
1203
1204fn full_piv_lu_meta(
1205    dtype: DType,
1206    shape: &[SymDim],
1207) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1208    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.full_piv_lu", shape)?;
1209    Ok(vec![
1210        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1211        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1212        (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1213        (dtype, matrix_shape(n.clone(), n, batch)),
1214        (singular_values_dtype(dtype), batch.to_vec()),
1215    ])
1216}
1217
1218fn svd_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1219    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd", shape)?;
1220    let k = m.clone().min(n.clone());
1221    Ok(vec![
1222        (dtype, matrix_shape(m, k.clone(), batch)),
1223        (singular_values_dtype(dtype), vector_shape(k.clone(), batch)),
1224        (dtype, matrix_shape(k, n, batch)),
1225    ])
1226}
1227
1228fn svd_full_meta(
1229    dtype: DType,
1230    shape: &[SymDim],
1231) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1232    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_full", shape)?;
1233    let k = m.clone().min(n.clone());
1234    Ok(vec![
1235        (dtype, matrix_shape(m.clone(), m, batch)),
1236        (singular_values_dtype(dtype), vector_shape(k, batch)),
1237        (dtype, matrix_shape(n.clone(), n, batch)),
1238    ])
1239}
1240
1241fn svd_values_meta(
1242    dtype: DType,
1243    shape: &[SymDim],
1244) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1245    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_values", shape)?;
1246    let k = m.min(n);
1247    Ok((singular_values_dtype(dtype), vector_shape(k, batch)))
1248}
1249
1250fn qr_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1251    let (m, n, batch) = matrix_meta_parts("tenferro-linalg.qr", shape)?;
1252    let k = m.clone().min(n.clone());
1253    Ok(vec![
1254        (dtype, matrix_shape(m, k.clone(), batch)),
1255        (dtype, matrix_shape(k, n, batch)),
1256    ])
1257}
1258
1259fn eigh_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1260    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh", shape)?;
1261    Ok(vec![
1262        (singular_values_dtype(dtype), vector_shape(n.clone(), batch)),
1263        (dtype, matrix_shape(n.clone(), n, batch)),
1264    ])
1265}
1266
1267fn eigh_values_meta(
1268    dtype: DType,
1269    shape: &[SymDim],
1270) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1271    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh_values", shape)?;
1272    Ok((singular_values_dtype(dtype), vector_shape(n, batch)))
1273}
1274
1275fn eig_meta(
1276    input_dtype: DType,
1277    shape: &[SymDim],
1278) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1279    let dtype = eig_output_dtype(input_dtype);
1280    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig", shape)?;
1281    Ok(vec![
1282        (dtype, vector_shape(n.clone(), batch)),
1283        (dtype, matrix_shape(n.clone(), n, batch)),
1284    ])
1285}
1286
1287fn eig_values_meta(
1288    input_dtype: DType,
1289    shape: &[SymDim],
1290) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1291    let dtype = eig_output_dtype(input_dtype);
1292    let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig_values", shape)?;
1293    Ok((dtype, vector_shape(n, batch)))
1294}
1295
1296fn matrix_shape(rows: SymDim, cols: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
1297    let mut shape = vec![rows, cols];
1298    shape.extend_from_slice(batch);
1299    shape
1300}
1301
1302fn vector_shape(len: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
1303    let mut shape = vec![len];
1304    shape.extend_from_slice(batch);
1305    shape
1306}
1307
1308fn eig_output_dtype(dtype: DType) -> DType {
1309    match dtype {
1310        DType::F64 | DType::C64 => DType::C64,
1311        DType::F32 | DType::C32 => DType::C32,
1312        DType::I32 | DType::I64 | DType::Bool => DType::C64,
1313    }
1314}
1315
1316fn singular_values_dtype(dtype: DType) -> DType {
1317    match dtype {
1318        DType::C64 => DType::F64,
1319        DType::C32 => DType::F32,
1320        other => other,
1321    }
1322}
1323
1324fn promote_dtypes(dtypes: &[DType]) -> DType {
1325    dtypes
1326        .iter()
1327        .copied()
1328        .reduce(tenferro_tensor::validate::promote_dtype)
1329        .unwrap_or(DType::F64)
1330}
1331
1332fn hash_dtype(hasher: &mut dyn Hasher, dtype: DType) {
1333    let tag = match dtype {
1334        DType::F64 => 0,
1335        DType::F32 => 1,
1336        DType::I64 => 2,
1337        DType::C64 => 3,
1338        DType::C32 => 4,
1339        DType::I32 => 5,
1340        DType::Bool => 6,
1341    };
1342    hasher.write_u8(tag);
1343}
1344
1345fn hash_svd_gauge(hasher: &mut dyn Hasher, gauge: SvdGauge) {
1346    let tag = match gauge {
1347        SvdGauge::Raw => 0,
1348        SvdGauge::CanonicalPivot => 1,
1349    };
1350    hasher.write_u8(tag);
1351}
1352
1353fn hash_eigh_gauge(hasher: &mut dyn Hasher, gauge: EighGauge) {
1354    let tag = match gauge {
1355        EighGauge::Raw => 0,
1356        EighGauge::CanonicalPivot => 1,
1357    };
1358    hasher.write_u8(tag);
1359}
1360
1361fn hash_qr_gauge(hasher: &mut dyn Hasher, gauge: QrGauge) {
1362    let tag = match gauge {
1363        QrGauge::Raw => 0,
1364        QrGauge::PositiveDiagonal => 1,
1365    };
1366    hasher.write_u8(tag);
1367}