Skip to main content

tenferro_linalg/
eager_ext.rs

1// Solve residual policy reference: PyTorch 8dd3b763, derivatives.yaml's
2// _linalg_solve_ex and FunctionsManual.cpp::linalg_solve_backward. The tracked
3// implementation runs the fused LuFactorSolve op, whose saved factors feed the
4// LuSolvePrepared adjoint solve.
5use std::sync::Arc;
6
7use tenferro_ad::error::{Error, Result};
8use tenferro_ad::extension::{
9    apply_eager_with_targeted_extension_in_session, apply_eager_with_targeted_extension_session,
10    EagerExtensionBackendKind, EagerExtensionTarget,
11};
12use tenferro_ad::{EagerSession, EagerTensor};
13use tenferro_cpu::CpuBackend;
14#[cfg(feature = "cuda")]
15use tenferro_gpu::cuda::CudaBackend;
16#[cfg(feature = "webgpu")]
17use tenferro_gpu::webgpu::WebGpuBackend;
18use tenferro_runtime::{ErrorPhase, ExtensionModule};
19
20use crate::eager_composites;
21use crate::extension::{
22    extension_module, validate_derivative_eps, EighOptions, LinalgExtensionOp, LinalgOp, QrOptions,
23    SvdOptions,
24};
25use crate::rank_revealing_qr::validate_rank_revealing_qr_options;
26use crate::{RankRevealingQrOptions, RankRevealingQrResult};
27
28/// Tensor-owned linear solve. Other eager linear-algebra operations use
29/// [`EagerSessionLinalgExt`] inside [`tenferro_ad::EagerRuntime::with_eager_session`].
30///
31/// The tensor-owned entry preserves calling-thread `no_grad` behavior for
32/// tracked solves; it must not be called from inside a borrowed session.
33///
34/// # Examples
35/// ```rust
36/// # use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
37/// # use tenferro_linalg::EagerTensorLinalgExt;
38/// # let ctx = EagerRuntime::new()?;
39/// # let a = EagerTensor::from_tensor_in(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?, ctx.clone())?;
40/// # let b = EagerTensor::from_tensor_in(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?, ctx)?;
41/// let x = a.solve(&b)?;
42/// assert_eq!(x.shape(), &[1, 1]);
43/// # Ok::<(), tenferro_ad::Error>(())
44/// ```
45#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
46pub trait EagerTensorLinalgExt {
47    /// # Errors
48    ///
49    /// Returns `Error::Validation` for incompatible matrix, batch, or dtype
50    /// metadata, `Error::Extension` for an unsupported dtype or singular
51    /// system, and `Error::RuntimeState` when the backend is unavailable.
52    /// # Examples
53    ///
54    /// ```rust
55    /// # use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
56    /// # use tenferro_cpu::CpuBackend;
57    /// # use tenferro_linalg::EagerTensorLinalgExt;
58    /// # let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
59    /// # let a = EagerTensor::from_tensor_in(
60    /// #     Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0]).unwrap(),
61    /// #     ctx.clone(),
62    /// # )?;
63    /// # let b = EagerTensor::from_tensor_in(
64    /// #     Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 8.0]).unwrap(),
65    /// #     ctx,
66    /// # )?;
67    /// let x = a.solve(&b)?;
68    /// assert_eq!(x.value()?.as_slice::<f64>()?, &[2.0, 2.0]);
69    /// # Ok::<(), tenferro_ad::Error>(())
70    /// ```
71    fn solve(&self, b: &EagerTensor) -> Result<EagerTensor>;
72}
73
74impl EagerTensorLinalgExt for EagerTensor {
75    fn solve(&self, b: &EagerTensor) -> Result<EagerTensor> {
76        solve(self, b)
77    }
78}
79
80/// Linear algebra operations on a runtime-bound borrowed eager session.
81///
82/// # Examples
83/// ```rust
84/// use tenferro_ad::{EagerRuntime, Tensor};
85/// use tenferro_linalg::EagerSessionLinalgExt;
86/// let ctx = EagerRuntime::new()?;
87/// let factor = ctx.with_eager_session(|session| {
88///     let input = session.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![4.0_f64])?)?;
89///     session.cholesky(&input)
90/// })?;
91/// assert_eq!(factor.value()?.as_slice::<f64>()?, &[2.0]);
92/// # Ok::<(), tenferro_ad::Error>(())
93/// ```
94#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
95pub trait EagerSessionLinalgExt {
96    /// Compute the lower Cholesky factor without reopening the eager backend.
97    ///
98    /// # Examples
99    /// ```rust
100    /// use tenferro_ad::{EagerRuntime, Tensor};
101    /// use tenferro_linalg::EagerSessionLinalgExt;
102    /// let ctx = EagerRuntime::new()?;
103    /// let factor = ctx.with_eager_session(|session| {
104    ///     let input = session.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![9.0_f64])?)?;
105    ///     session.cholesky(&input)
106    /// })?;
107    /// assert_eq!(factor.value()?.as_slice::<f64>()?, &[3.0]);
108    /// # Ok::<(), tenferro_ad::Error>(())
109    /// ```
110    /// # Errors
111    /// Returns typed validation, extension, backend, module or unsupported-executor errors.
112    fn cholesky(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
113
114    /// Compute a thin singular value decomposition in this borrowed session.
115    ///
116    /// # Examples
117    /// ```rust
118    /// use tenferro_ad::{EagerRuntime, Tensor};
119    /// use tenferro_linalg::EagerSessionLinalgExt;
120    /// let ctx = EagerRuntime::new()?;
121    /// let (_u, values, _vt) = ctx.with_eager_session(|s| {
122    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![3.0_f64])?)?;
123    ///     s.svd(&a)
124    /// })?;
125    /// assert_eq!(values.value()?.as_slice::<f64>()?, &[3.0]);
126    /// # Ok::<(), tenferro_ad::Error>(())
127    /// ```
128    /// # Errors
129    /// Returns typed validation, extension, backend, or unsupported-executor errors.
130    fn svd(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor)>;
131
132    /// Compute thin SVD with an explicit gauge, driver, and derivative regularizer.
133    ///
134    /// # Examples
135    /// ```rust
136    /// use tenferro_ad::{EagerRuntime, Tensor};
137    /// use tenferro_linalg::{EagerSessionLinalgExt, SvdOptions};
138    /// let ctx = EagerRuntime::new()?;
139    /// let (_u, values, _vt) = ctx.with_eager_session(|s| {
140    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![3.0_f64])?)?;
141    ///     s.svd_with_options(&a, SvdOptions::default())
142    /// })?;
143    /// assert_eq!(values.value()?.as_slice::<f64>()?, &[3.0]);
144    /// # Ok::<(), tenferro_ad::Error>(())
145    /// ```
146    /// # Errors
147    /// Returns typed invalid-tolerance, extension, backend, or unsupported errors.
148    fn svd_with_options(
149        &mut self,
150        input: &EagerTensor,
151        options: SvdOptions,
152    ) -> Result<(EagerTensor, EagerTensor, EagerTensor)>;
153
154    /// Compute full-matrices SVD, including the right nullspace.
155    ///
156    /// # Examples
157    /// ```rust
158    /// use tenferro_ad::{EagerRuntime, Tensor};
159    /// use tenferro_linalg::EagerSessionLinalgExt;
160    /// let ctx = EagerRuntime::new()?;
161    /// let (u, values, vh) = ctx.with_eager_session(|s| {
162    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 2], vec![1.0_f64, 1.0])?)?;
163    ///     s.svd_full(&a)
164    /// })?;
165    /// assert_eq!(u.shape(), &[1, 1]);
166    /// assert_eq!(values.shape(), &[1]);
167    /// assert_eq!(vh.shape(), &[2, 2]);
168    /// # Ok::<(), tenferro_ad::Error>(())
169    /// ```
170    /// # Errors
171    /// Returns `Error::Validation` with `ValidationError::RankMismatch` when the
172    /// input is not a (batched) matrix, `Error::UnsupportedAdRule` when a
173    /// traced input needs a derivative this decomposition does not provide,
174    /// `Error::Extension` carrying the linalg failure
175    /// (for example `Error::NonConvergence` or an unsupported dtype), or
176    /// `Error::TensorRuntime` for a backend failure.
177    fn svd_full(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor)>;
178
179    /// Compute a thin QR decomposition in this borrowed session.
180    ///
181    /// # Examples
182    /// ```rust
183    /// use tenferro_ad::{EagerRuntime, Tensor};
184    /// use tenferro_linalg::EagerSessionLinalgExt;
185    /// let ctx = EagerRuntime::new()?;
186    /// let (q, r) = ctx.with_eager_session(|s| {
187    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![3.0_f64])?)?;
188    ///     s.qr(&a)
189    /// })?;
190    /// assert_eq!(q.shape(), &[1, 1]);
191    /// assert_eq!(r.shape(), &[1, 1]);
192    /// # Ok::<(), tenferro_ad::Error>(())
193    /// ```
194    /// # Errors
195    /// Returns typed validation, extension, backend, or unsupported-executor errors.
196    fn qr(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)>;
197
198    /// Compute QR with an explicit post-processing gauge.
199    ///
200    /// # Examples
201    /// ```rust
202    /// use tenferro_ad::{EagerRuntime, Tensor};
203    /// use tenferro_linalg::{EagerSessionLinalgExt, QrOptions};
204    /// let ctx = EagerRuntime::new()?;
205    /// let (q, r) = ctx.with_eager_session(|s| {
206    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![3.0_f64])?)?;
207    ///     s.qr_with_options(&a, QrOptions::default())
208    /// })?;
209    /// assert_eq!(q.shape(), &[1, 1]);
210    /// assert_eq!(r.shape(), &[1, 1]);
211    /// # Ok::<(), tenferro_ad::Error>(())
212    /// ```
213    /// # Errors
214    /// Returns typed validation, extension, backend, or unsupported errors.
215    fn qr_with_options(
216        &mut self,
217        input: &EagerTensor,
218        options: QrOptions,
219    ) -> Result<(EagerTensor, EagerTensor)>;
220
221    /// Solve a triangular system in this borrowed session.
222    ///
223    /// # Examples
224    /// ```rust
225    /// use tenferro_ad::{EagerRuntime, Tensor};
226    /// use tenferro_linalg::EagerSessionLinalgExt;
227    /// let ctx = EagerRuntime::new()?;
228    /// let x = ctx.with_eager_session(|s| {
229    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
230    ///     let b = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
231    ///     s.triangular_solve(&a, &b, true, true, false, false)
232    /// })?;
233    /// assert_eq!(x.value()?.as_slice::<f64>()?, &[2.0]);
234    /// # Ok::<(), tenferro_ad::Error>(())
235    /// ```
236    /// # Errors
237    /// Returns typed validation, extension, backend, or unsupported-executor errors.
238    fn triangular_solve(
239        &mut self,
240        matrix: &EagerTensor,
241        rhs: &EagerTensor,
242        left_side: bool,
243        lower: bool,
244        transpose_a: bool,
245        unit_diagonal: bool,
246    ) -> Result<EagerTensor>;
247
248    /// Solve a square linear system in this borrowed session.
249    ///
250    /// # Examples
251    /// ```rust
252    /// use tenferro_ad::{EagerRuntime, Tensor};
253    /// use tenferro_linalg::EagerSessionLinalgExt;
254    /// let ctx = EagerRuntime::new()?;
255    /// let x = ctx.with_eager_session(|s| {
256    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
257    ///     let b = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
258    ///     s.solve(&a, &b)
259    /// })?;
260    /// assert_eq!(x.value()?.as_slice::<f64>()?, &[2.0]);
261    /// # Ok::<(), tenferro_ad::Error>(())
262    /// ```
263    /// # Errors
264    /// Returns typed validation, extension, backend, or unsupported-executor errors.
265    fn solve(&mut self, matrix: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor>;
266
267    /// Solve a full-column-rank least-squares problem inside this session.
268    ///
269    /// # Examples
270    /// ```rust
271    /// use tenferro_ad::{EagerRuntime, Tensor};
272    /// use tenferro_linalg::EagerSessionLinalgExt;
273    /// let ctx = EagerRuntime::new()?;
274    /// let x = ctx.with_eager_session(|s| {
275    ///     let a = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![1.0_f64, 2.0])?)?;
276    ///     let b = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![2.0_f64, 4.0])?)?;
277    ///     s.lstsq(&a, &b)
278    /// })?;
279    /// assert!((x.value()?.as_slice::<f64>()?[0] - 2.0).abs() < 1e-12);
280    /// # Ok::<(), tenferro_ad::Error>(())
281    /// ```
282    /// # Errors
283    /// Returns typed rank/shape/dtype validation, extension, backend, or unsupported errors.
284    fn lstsq(&mut self, matrix: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor>;
285
286    /// Factor a matrix without leaving this borrowed session.
287    ///
288    /// # Examples
289    /// ```rust
290    /// use tenferro_ad::{EagerRuntime, Tensor};
291    /// use tenferro_linalg::EagerSessionLinalgExt;
292    /// let ctx = EagerRuntime::new()?;
293    /// let (_p, _l, u, _parity) = ctx.with_eager_session(|s| {
294    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
295    ///     s.lu(&a)
296    /// })?;
297    /// assert_eq!(u.value()?.as_slice::<f64>()?, &[2.0]);
298    /// # Ok::<(), tenferro_ad::Error>(())
299    /// ```
300    /// # Errors
301    /// Returns typed validation, extension, backend, or unsupported-executor errors.
302    fn lu(
303        &mut self,
304        input: &EagerTensor,
305    ) -> Result<(EagerTensor, EagerTensor, EagerTensor, EagerTensor)>;
306
307    /// Compute determinant sign and logarithm of its absolute value in this session.
308    ///
309    /// # Examples
310    /// ```rust
311    /// use tenferro_ad::{EagerRuntime, Tensor};
312    /// use tenferro_linalg::EagerSessionLinalgExt;
313    /// let ctx = EagerRuntime::new()?;
314    /// let (sign, logabs) = ctx.with_eager_session(|s| {
315    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
316    ///     s.slogdet(&a)
317    /// })?;
318    /// assert_eq!(sign.value()?.as_slice::<f64>()?, &[1.0]);
319    /// assert!((logabs.value()?.as_slice::<f64>()?[0] - 2.0_f64.ln()).abs() < 1e-12);
320    /// # Ok::<(), tenferro_ad::Error>(())
321    /// ```
322    /// # Errors
323    /// Returns typed validation, extension, backend, or unsupported-executor errors.
324    fn slogdet(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)>;
325
326    /// Compute a determinant inside the borrowed session.
327    ///
328    /// # Examples
329    /// ```rust
330    /// use tenferro_ad::{EagerRuntime, Tensor};
331    /// use tenferro_linalg::EagerSessionLinalgExt;
332    /// let ctx = EagerRuntime::new()?;
333    /// let det = ctx.with_eager_session(|s| {
334    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
335    ///     s.det(&a)
336    /// })?;
337    /// assert_eq!(det.value()?.as_slice::<f64>()?, &[2.0]);
338    /// # Ok::<(), tenferro_ad::Error>(())
339    /// ```
340    /// # Errors
341    /// Returns typed validation, extension, backend, or unsupported-executor errors.
342    fn det(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
343
344    /// Invert a square matrix inside this borrowed session.
345    ///
346    /// # Examples
347    /// ```rust
348    /// use tenferro_ad::{EagerRuntime, Tensor};
349    /// use tenferro_linalg::EagerSessionLinalgExt;
350    /// let ctx = EagerRuntime::new()?;
351    /// let inverse = ctx.with_eager_session(|s| {
352    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
353    ///     s.inv(&a)
354    /// })?;
355    /// assert_eq!(inverse.value()?.as_slice::<f64>()?, &[0.5]);
356    /// # Ok::<(), tenferro_ad::Error>(())
357    /// ```
358    /// # Errors
359    /// Returns typed rank/shape/dtype validation, extension, backend, or unsupported errors.
360    fn inv(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
361
362    /// Compute eigenvalues of a Hermitian matrix in this borrowed session.
363    ///
364    /// # Examples
365    /// ```rust
366    /// use tenferro_ad::{EagerRuntime, Tensor};
367    /// use tenferro_linalg::EagerSessionLinalgExt;
368    /// let ctx = EagerRuntime::new()?;
369    /// let values = ctx.with_eager_session(|s| {
370    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
371    ///     s.eigvalsh(&a)
372    /// })?;
373    /// assert_eq!(values.value()?.as_slice::<f64>()?, &[4.0]);
374    /// # Ok::<(), tenferro_ad::Error>(())
375    /// ```
376    /// # Errors
377    /// Returns typed validation, extension, backend, or unsupported-executor errors.
378    fn eigvalsh(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
379
380    /// Compute general eigenvalues in this borrowed session.
381    ///
382    /// # Examples
383    /// ```rust
384    /// use tenferro_ad::{EagerRuntime, Tensor};
385    /// use tenferro_linalg::EagerSessionLinalgExt;
386    /// let ctx = EagerRuntime::new()?;
387    /// let values = ctx.with_eager_session(|s| {
388    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
389    ///     s.eigvals(&a)
390    /// })?;
391    /// assert_eq!(values.value()?.as_slice::<num_complex::Complex64>()?, &[num_complex::Complex64::new(4.0, 0.0)]);
392    /// # Ok::<(), tenferro_ad::Error>(())
393    /// ```
394    /// # Errors
395    /// Returns typed validation, extension, backend, or unsupported-executor errors.
396    fn eigvals(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
397
398    /// Compute eigenvalues and vectors of a Hermitian matrix in this session.
399    ///
400    /// # Examples
401    /// ```rust
402    /// use tenferro_ad::{EagerRuntime, Tensor};
403    /// use tenferro_linalg::EagerSessionLinalgExt;
404    /// let ctx = EagerRuntime::new()?;
405    /// let (values, vectors) = ctx.with_eager_session(|s| {
406    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
407    ///     s.eigh(&a)
408    /// })?;
409    /// assert_eq!(values.value()?.as_slice::<f64>()?, &[4.0]);
410    /// assert_eq!(vectors.shape(), &[1, 1]);
411    /// # Ok::<(), tenferro_ad::Error>(())
412    /// ```
413    /// # Errors
414    /// Returns typed validation, unsupported-dtype, extension, or backend errors.
415    fn eigh(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)>;
416
417    /// Compute Hermitian eigendecomposition with explicit gauge and tolerance.
418    ///
419    /// # Examples
420    /// ```rust
421    /// use tenferro_ad::{EagerRuntime, Tensor};
422    /// use tenferro_linalg::{EagerSessionLinalgExt, EighOptions};
423    /// let ctx = EagerRuntime::new()?;
424    /// let (values, _vectors) = ctx.with_eager_session(|s| {
425    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
426    ///     s.eigh_with_options(&a, EighOptions::default())
427    /// })?;
428    /// assert_eq!(values.value()?.as_slice::<f64>()?, &[4.0]);
429    /// # Ok::<(), tenferro_ad::Error>(())
430    /// ```
431    /// # Errors
432    /// Returns `Error::Validation` with `ValidationError::InvalidArgument` for an
433    /// invalid tolerance or with `ValidationError::ShapeMismatch` for a
434    /// non-square input, `Error::Extension` carrying the linalg failure
435    /// (for example `Error::NonConvergence` or an unsupported dtype), or
436    /// `Error::TensorRuntime` for a backend failure.
437    fn eigh_with_options(
438        &mut self,
439        input: &EagerTensor,
440        options: EighOptions,
441    ) -> Result<(EagerTensor, EagerTensor)>;
442
443    /// Compute general eigenvalues and eigenvectors inside this session.
444    ///
445    /// # Examples
446    /// ```rust
447    /// use tenferro_ad::{EagerRuntime, Tensor};
448    /// use tenferro_linalg::EagerSessionLinalgExt;
449    /// let ctx = EagerRuntime::new()?;
450    /// let (values, vectors) = ctx.with_eager_session(|s| {
451    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
452    ///     s.eig(&a)
453    /// })?;
454    /// assert_eq!(values.value()?.as_slice::<num_complex::Complex64>()?, &[num_complex::Complex64::new(4.0, 0.0)]);
455    /// assert_eq!(vectors.shape(), &[1, 1]);
456    /// # Ok::<(), tenferro_ad::Error>(())
457    /// ```
458    /// # Errors
459    /// Returns `Error::Validation` with `ValidationError::RankMismatch` or
460    /// `ValidationError::ShapeMismatch` for a non-square (batched) matrix,
461    /// `Error::Extension` carrying the linalg failure
462    /// (for example `Error::NonConvergence` or an unsupported dtype), or
463    /// `Error::TensorRuntime` for a backend failure.
464    fn eig(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)>;
465
466    /// Compute the Moore-Penrose pseudoinverse using the default tolerance.
467    ///
468    /// # Examples
469    /// ```rust
470    /// use tenferro_ad::{EagerRuntime, Tensor};
471    /// use tenferro_linalg::EagerSessionLinalgExt;
472    /// let ctx = EagerRuntime::new()?;
473    /// let inverse = ctx.with_eager_session(|s| {
474    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
475    ///     s.pinv(&a)
476    /// })?;
477    /// assert_eq!(inverse.value()?.as_slice::<f64>()?, &[0.5]);
478    /// # Ok::<(), tenferro_ad::Error>(())
479    /// ```
480    /// # Errors
481    /// Returns `Error::Validation` with `ValidationError::RankMismatch` when the
482    /// input is not a (batched) matrix, `Error::Extension` carrying the linalg failure
483    /// (for example `Error::NonConvergence` or an unsupported dtype) from the
484    /// underlying SVD, or `Error::TensorRuntime` for a backend failure.
485    fn pinv(&mut self, input: &EagerTensor) -> Result<EagerTensor>;
486
487    /// Compute the pseudoinverse with an explicit relative tolerance.
488    ///
489    /// # Examples
490    /// ```rust
491    /// use tenferro_ad::{EagerRuntime, Tensor};
492    /// use tenferro_linalg::EagerSessionLinalgExt;
493    /// let ctx = EagerRuntime::new()?;
494    /// let inverse = ctx.with_eager_session(|s| {
495    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
496    ///     s.pinv_with_rtol(&a, 1.0e-12)
497    /// })?;
498    /// assert_eq!(inverse.value()?.as_slice::<f64>()?, &[0.5]);
499    /// # Ok::<(), tenferro_ad::Error>(())
500    /// ```
501    /// # Errors
502    /// Returns `Error::Validation` with `ValidationError::RankMismatch` when the
503    /// input is not a (batched) matrix, `ValidationError::InvalidArgument` for a
504    /// negative or non-finite tolerance, `Error::Extension` carrying the linalg failure
505    /// (for example `Error::NonConvergence` or an unsupported dtype) from
506    /// the underlying SVD, or `Error::TensorRuntime` for a backend failure.
507    fn pinv_with_rtol(&mut self, input: &EagerTensor, rtol: f64) -> Result<EagerTensor>;
508
509    /// Compute a vector, matrix, or tensor norm in this session.
510    /// An empty axis list is a no-op and clones the input without dispatch.
511    ///
512    /// # Examples
513    /// ```rust
514    /// use tenferro_ad::{EagerRuntime, Tensor};
515    /// use tenferro_linalg::EagerSessionLinalgExt;
516    /// let ctx = EagerRuntime::new()?;
517    /// let result = ctx.with_eager_session(|s| {
518    ///     let a = s.constant_from(Tensor::from_vec_col_major([2], vec![3.0_f64, 4.0])?)?;
519    ///     s.norm(&a, Some(2.0), Some(&[0]), false)
520    /// })?;
521    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[5.0]);
522    /// # Ok::<(), tenferro_ad::Error>(())
523    /// ```
524    /// # Errors
525    /// Returns typed unsupported-dtype, invalid-axis/order, SVD, or backend errors.
526    fn norm(
527        &mut self,
528        input: &EagerTensor,
529        ord: Option<f64>,
530        dim: Option<&[usize]>,
531        keepdim: bool,
532    ) -> Result<EagerTensor>;
533
534    /// Compute complete-pivot LU factors `(P, L, U, Q, parity)` in this session.
535    /// Reconstruction uses `A = P^T * L * U * Q`; scalar `parity` is real
536    /// (`F32` for `F32`/`C32` inputs and `F64` for `F64`/`C64`).
537    ///
538    /// # Examples
539    /// ```rust
540    /// use tenferro_ad::{EagerRuntime, Tensor};
541    /// use tenferro_linalg::EagerSessionLinalgExt;
542    /// let ctx = EagerRuntime::new()?;
543    /// let (p, _l, _u, q, parity) = ctx.with_eager_session(|s| {
544    ///     let a = s.constant_from(Tensor::from_vec_col_major([2, 2], vec![1.0_f64, 3.0, 2.0, 4.0])?)?;
545    ///     s.full_piv_lu(&a)
546    /// })?;
547    /// assert_eq!(p.shape(), &[2, 2]);
548    /// assert_eq!(q.shape(), &[2, 2]);
549    /// assert_eq!(parity.shape(), &[] as &[usize]);
550    /// # Ok::<(), tenferro_ad::Error>(())
551    /// ```
552    /// # Errors
553    /// Returns typed rank/shape, unsupported-provider, numerical, or output-count errors.
554    fn full_piv_lu(
555        &mut self,
556        input: &EagerTensor,
557    ) -> Result<(
558        EagerTensor,
559        EagerTensor,
560        EagerTensor,
561        EagerTensor,
562        EagerTensor,
563    )>;
564
565    /// Solve a linear system using complete-pivot LU in this session.
566    ///
567    /// # Examples
568    /// ```rust
569    /// use tenferro_ad::{EagerRuntime, Tensor};
570    /// use tenferro_linalg::EagerSessionLinalgExt;
571    /// let ctx = EagerRuntime::new()?;
572    /// let x = ctx.with_eager_session(|s| {
573    ///     let a = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
574    ///     let b = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![4.0_f64])?)?;
575    ///     s.full_piv_lu_solve(&a, &b)
576    /// })?;
577    /// assert_eq!(x.value()?.as_slice::<f64>()?, &[2.0]);
578    /// # Ok::<(), tenferro_ad::Error>(())
579    /// ```
580    /// # Errors
581    /// Returns typed rank/shape, unsupported-provider, singularity, or backend errors.
582    fn full_piv_lu_solve(&mut self, matrix: &EagerTensor, rhs: &EagerTensor)
583        -> Result<EagerTensor>;
584
585    /// Compute column-pivoted rank-revealing QR in this session.
586    ///
587    /// # Examples
588    /// ```rust
589    /// use tenferro_ad::{EagerRuntime, Tensor};
590    /// use tenferro_linalg::{EagerSessionLinalgExt, RankRevealingQrOptions};
591    /// let ctx = EagerRuntime::new()?;
592    /// let result = ctx.with_eager_session(|s| {
593    ///     let a = s.constant_from(Tensor::from_vec_col_major([2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?)?;
594    ///     s.rank_revealing_qr(&a, RankRevealingQrOptions::default())
595    /// })?;
596    /// assert_eq!(result.column_permutation.shape(), &[2]);
597    /// assert_eq!(result.rank.value()?.as_slice::<i64>()?, &[2]);
598    /// # Ok::<(), tenferro_ad::Error>(())
599    /// ```
600    /// # Errors
601    /// Returns `Error::Validation` with `ValidationError::RankMismatch`,
602    /// `ValidationError::DTypeMismatch` or `ValidationError::InvalidArgument` for
603    /// an invalid rank, dtype or tolerance, `Error::Extension` carrying the linalg failure
604    /// (for example `Error::NonConvergence` or an unsupported dtype), or
605    /// `Error::TensorRuntime` for a backend failure.
606    fn rank_revealing_qr(
607        &mut self,
608        input: &EagerTensor,
609        options: RankRevealingQrOptions,
610    ) -> Result<RankRevealingQrResult<EagerTensor>>;
611
612    /// Initialize compact Householder QR state in this borrowed session.
613    ///
614    /// # Examples
615    /// ```rust
616    /// use tenferro_ad::{EagerRuntime, Tensor};
617    /// use tenferro_linalg::EagerSessionLinalgExt;
618    /// let ctx = EagerRuntime::new()?;
619    /// let state = ctx.with_eager_session(|s| {
620    ///     let a = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![1.0_f64, 2.0])?)?;
621    ///     s.householder_qr(&a)
622    /// })?;
623    /// assert!(format!("{state:?}").starts_with("HouseholderQr"));
624    /// # Ok::<(), tenferro_ad::Error>(())
625    /// ```
626    /// # Errors
627    /// Returns typed invalid metadata, unsupported executor, or backend errors.
628    fn householder_qr(&mut self, input: &EagerTensor) -> Result<crate::HouseholderQr<EagerTensor>>;
629}
630
631impl EagerSessionLinalgExt for EagerSession<'_> {
632    fn svd(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
633        self.svd_with_options(input, SvdOptions::default())
634    }
635
636    fn svd_with_options(
637        &mut self,
638        input: &EagerTensor,
639        options: SvdOptions,
640    ) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
641        validate_derivative_eps("svd_with_options", options.derivative_eps)?;
642        three_outputs(
643            apply_linalg_eager_in_session(
644                self,
645                LinalgOp::Svd {
646                    derivative_eps: options.derivative_eps,
647                    gauge: options.gauge,
648                    driver: options.driver,
649                },
650                &[input],
651            )?,
652            "svd",
653        )
654    }
655
656    fn svd_full(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
657        three_outputs(
658            apply_linalg_eager_in_session(self, LinalgOp::SvdFull, &[input])?,
659            "svd_full",
660        )
661    }
662
663    fn qr(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
664        self.qr_with_options(input, QrOptions::default())
665    }
666
667    fn qr_with_options(
668        &mut self,
669        input: &EagerTensor,
670        options: QrOptions,
671    ) -> Result<(EagerTensor, EagerTensor)> {
672        two_outputs(
673            apply_linalg_eager_in_session(
674                self,
675                LinalgOp::Qr {
676                    gauge: options.gauge,
677                },
678                &[input],
679            )?,
680            "qr",
681        )
682    }
683
684    fn cholesky(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
685        one_output(
686            apply_linalg_eager_in_session(self, LinalgOp::Cholesky, &[input])?,
687            "cholesky",
688        )
689    }
690
691    fn triangular_solve(
692        &mut self,
693        matrix: &EagerTensor,
694        rhs: &EagerTensor,
695        left_side: bool,
696        lower: bool,
697        transpose_a: bool,
698        unit_diagonal: bool,
699    ) -> Result<EagerTensor> {
700        one_output(
701            apply_linalg_eager_in_session(
702                self,
703                LinalgOp::TriangularSolve {
704                    left_side,
705                    lower,
706                    transpose_a,
707                    unit_diagonal,
708                },
709                &[matrix, rhs],
710            )?,
711            "triangular_solve",
712        )
713    }
714
715    fn solve(&mut self, matrix: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
716        if !matrix.tracks_grad() && !rhs.tracks_grad() {
717            return one_output(
718                apply_linalg_eager_in_session(self, LinalgOp::Solve, &[matrix, rhs])?,
719                "solve",
720            );
721        }
722        validate_tracked_solve_inputs(matrix, rhs)?;
723        factor_solve_output(apply_linalg_eager_in_session(
724            self,
725            LinalgOp::LuFactorSolve,
726            &[matrix, rhs],
727        )?)
728    }
729
730    fn lstsq(&mut self, matrix: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
731        eager_composites::lstsq(self, matrix, rhs)
732    }
733
734    fn lu(
735        &mut self,
736        input: &EagerTensor,
737    ) -> Result<(EagerTensor, EagerTensor, EagerTensor, EagerTensor)> {
738        let mut outputs = apply_linalg_eager_in_session(self, LinalgOp::Lu, &[input])?.into_iter();
739        match (
740            outputs.next(),
741            outputs.next(),
742            outputs.next(),
743            outputs.next(),
744            outputs.next(),
745        ) {
746            (Some(p), Some(l), Some(u), Some(parity), None) => Ok((p, l, u, parity)),
747            _ => Err(Error::Internal(
748                "lu eager op returned an unexpected number of outputs".into(),
749            )),
750        }
751    }
752
753    fn slogdet(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
754        eager_composites::slogdet(self, input)
755    }
756
757    fn det(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
758        eager_composites::det(self, input)
759    }
760
761    fn inv(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
762        eager_composites::inv(self, input)
763    }
764
765    fn eigvalsh(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
766        one_output(
767            apply_linalg_eager_in_session(
768                self,
769                LinalgOp::EighVals {
770                    derivative_eps: crate::extension::DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
771                    driver: EighOptions::default().driver,
772                },
773                &[input],
774            )?,
775            "eigvalsh",
776        )
777    }
778
779    fn eigvals(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
780        one_output(
781            apply_linalg_eager_in_session(
782                self,
783                LinalgOp::EigVals {
784                    input_dtype: input.dtype(),
785                },
786                &[input],
787            )?,
788            "eigvals",
789        )
790    }
791
792    fn eigh(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
793        self.eigh_with_options(input, EighOptions::default())
794    }
795
796    fn eigh_with_options(
797        &mut self,
798        input: &EagerTensor,
799        options: EighOptions,
800    ) -> Result<(EagerTensor, EagerTensor)> {
801        validate_derivative_eps("eigh_with_options", options.derivative_eps)?;
802        two_outputs(
803            apply_linalg_eager_in_session(
804                self,
805                LinalgOp::Eigh {
806                    derivative_eps: options.derivative_eps,
807                    gauge: options.gauge,
808                    driver: options.driver,
809                },
810                &[input],
811            )?,
812            "eigh",
813        )
814    }
815
816    fn eig(&mut self, input: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
817        two_outputs(
818            apply_linalg_eager_in_session(
819                self,
820                LinalgOp::Eig {
821                    input_dtype: input.dtype(),
822                },
823                &[input],
824            )?,
825            "eig",
826        )
827    }
828
829    fn pinv(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
830        eager_composites::pinv(self, input)
831    }
832
833    fn pinv_with_rtol(&mut self, input: &EagerTensor, rtol: f64) -> Result<EagerTensor> {
834        eager_composites::pinv_with_rtol(self, input, rtol)
835    }
836
837    fn norm(
838        &mut self,
839        input: &EagerTensor,
840        ord: Option<f64>,
841        dim: Option<&[usize]>,
842        keepdim: bool,
843    ) -> Result<EagerTensor> {
844        eager_composites::norm(self, input, ord, dim, keepdim)
845    }
846
847    fn full_piv_lu(
848        &mut self,
849        input: &EagerTensor,
850    ) -> Result<(
851        EagerTensor,
852        EagerTensor,
853        EagerTensor,
854        EagerTensor,
855        EagerTensor,
856    )> {
857        let mut outputs =
858            apply_linalg_eager_in_session(self, LinalgOp::FullPivLu, &[input])?.into_iter();
859        match (
860            outputs.next(),
861            outputs.next(),
862            outputs.next(),
863            outputs.next(),
864            outputs.next(),
865            outputs.next(),
866        ) {
867            (Some(p), Some(l), Some(u), Some(q), Some(parity), None) => Ok((p, l, u, q, parity)),
868            _ => Err(Error::Internal(
869                "full_piv_lu eager op returned an unexpected number of outputs".into(),
870            )),
871        }
872    }
873
874    fn full_piv_lu_solve(
875        &mut self,
876        matrix: &EagerTensor,
877        rhs: &EagerTensor,
878    ) -> Result<EagerTensor> {
879        one_output(
880            apply_linalg_eager_in_session(
881                self,
882                LinalgOp::FullPivLuSolve { transpose_a: false },
883                &[matrix, rhs],
884            )?,
885            "full_piv_lu_solve",
886        )
887    }
888
889    fn rank_revealing_qr(
890        &mut self,
891        input: &EagerTensor,
892        options: RankRevealingQrOptions,
893    ) -> Result<RankRevealingQrResult<EagerTensor>> {
894        validate_rank_revealing_qr_options("rank_revealing_qr", options)?;
895        let mut outputs = apply_linalg_eager_in_session(
896            self,
897            LinalgOp::RankRevealingQr {
898                gauge: options.gauge,
899                rtol: options.rtol,
900                atol: options.atol,
901            },
902            &[input],
903        )?
904        .into_iter();
905        match (
906            outputs.next(),
907            outputs.next(),
908            outputs.next(),
909            outputs.next(),
910            outputs.next(),
911        ) {
912            (Some(q), Some(r), Some(column_permutation), Some(rank), None) => {
913                Ok(RankRevealingQrResult {
914                    q,
915                    r,
916                    column_permutation,
917                    rank,
918                })
919            }
920            _ => Err(Error::Internal(
921                "rank_revealing_qr eager op returned an unexpected number of outputs".into(),
922            )),
923        }
924    }
925
926    fn householder_qr(&mut self, input: &EagerTensor) -> Result<crate::HouseholderQr<EagerTensor>> {
927        crate::householder::eager_state(apply_linalg_eager_in_session(
928            self,
929            LinalgOp::HouseholderQrFactor,
930            &[input],
931        )?)
932    }
933}
934
935pub(crate) fn apply_linalg_eager_in_session(
936    session: &mut EagerSession<'_>,
937    op: LinalgOp,
938    inputs: &[&EagerTensor],
939) -> Result<Vec<EagerTensor>> {
940    let op = Arc::new(LinalgExtensionOp::new(op));
941    apply_eager_with_targeted_extension_in_session(session, op, inputs, eager_extension_module)
942}
943
944pub(crate) fn apply_linalg_eager(
945    op: LinalgOp,
946    inputs: &[&EagerTensor],
947) -> Result<Vec<EagerTensor>> {
948    let op = Arc::new(LinalgExtensionOp::new(op));
949    apply_eager_with_targeted_extension_session(op, inputs, eager_extension_module)
950}
951
952fn eager_extension_module(target: EagerExtensionTarget) -> Result<Arc<dyn ExtensionModule>> {
953    let EagerExtensionTarget {
954        engine_id,
955        backend_kind,
956    } = target;
957    match backend_kind {
958        EagerExtensionBackendKind::Cpu => {
959            extension_module::<CpuBackend>(engine_id).map_err(eager_runtime_config_error)
960        }
961        #[cfg(feature = "cuda")]
962        EagerExtensionBackendKind::Cuda => {
963            extension_module::<CudaBackend>(engine_id).map_err(eager_runtime_config_error)
964        }
965        #[cfg(feature = "webgpu")]
966        EagerExtensionBackendKind::WebGpu => {
967            extension_module::<WebGpuBackend>(engine_id).map_err(eager_runtime_config_error)
968        }
969    }
970}
971
972fn eager_runtime_config_error(source: tenferro_runtime::RuntimeConfigError) -> Error {
973    Error::runtime_state_source(
974        "tenferro_linalg::eager_extension_module",
975        ErrorPhase::Execution,
976        source,
977    )
978}
979
980/// Solve a linear system for eager tensors.
981///
982/// # Examples
983///
984/// ```rust
985/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
986/// use tenferro_linalg::EagerTensorLinalgExt;
987///
988/// let ctx = EagerRuntime::new()?;
989/// let a = EagerTensor::from_tensor_in(
990///     Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0]).unwrap(),
991///     ctx.clone(),
992/// ).unwrap();
993/// let b = EagerTensor::from_tensor_in(
994///     Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 8.0]).unwrap(),
995///     ctx,
996/// ).unwrap();
997/// let x = a.solve(&b)?;
998/// assert_eq!(x.shape(), &[2, 1]);
999/// # Ok::<(), tenferro_ad::Error>(())
1000/// ```
1001///
1002/// # Errors
1003///
1004/// Returns `Error::Validation` for incompatible matrix, batch, or dtype
1005/// metadata, `Error::Extension` for an unsupported dtype or singular system,
1006/// and `Error::RuntimeState` when the backend is unavailable.
1007pub fn solve(a: &EagerTensor, b: &EagerTensor) -> Result<EagerTensor> {
1008    if !a.tracks_grad() && !b.tracks_grad() {
1009        return one_output(apply_linalg_eager(LinalgOp::Solve, &[a, b])?, "solve");
1010    }
1011    validate_tracked_solve_inputs(a, b)?;
1012    factor_solve_output(apply_linalg_eager(LinalgOp::LuFactorSolve, &[a, b])?)
1013}
1014
1015fn validate_tracked_solve_inputs(a: &EagerTensor, b: &EagerTensor) -> Result<()> {
1016    if !a.same_context(b) {
1017        return Err(Error::ContextMismatch {
1018            lhs: a.ctx_id(),
1019            rhs: b.ctx_id(),
1020        });
1021    }
1022    crate::validation::validate_solve_inputs(a.dtype(), a.shape(), b.dtype(), b.shape())
1023}
1024
1025// One fused factor+solve retains LU/pivots and X for backward, while the
1026// explicit A operand preserves higher-order semantics (PyTorch's
1027// _linalg_solve_ex / FunctionsManual.cpp::linalg_solve_backward).
1028fn factor_solve_output(outputs: Vec<EagerTensor>) -> Result<EagerTensor> {
1029    let mut outputs = outputs.into_iter();
1030    match (
1031        outputs.next(),
1032        outputs.next(),
1033        outputs.next(),
1034        outputs.next(),
1035    ) {
1036        (Some(x), Some(_packed_lu), Some(_pivots), None) => Ok(x),
1037        _ => Err(Error::Internal(
1038            "lu_factor_solve eager op returned an unexpected number of outputs".into(),
1039        )),
1040    }
1041}
1042
1043fn three_outputs(
1044    outputs: Vec<EagerTensor>,
1045    name: &str,
1046) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
1047    let mut outputs = outputs.into_iter();
1048    match (
1049        outputs.next(),
1050        outputs.next(),
1051        outputs.next(),
1052        outputs.next(),
1053    ) {
1054        (Some(first), Some(second), Some(third), None) => Ok((first, second, third)),
1055        _ => Err(Error::Internal(format!(
1056            "{name} eager op returned an unexpected number of outputs"
1057        ))),
1058    }
1059}
1060
1061pub(crate) fn one_output(outputs: Vec<EagerTensor>, name: &str) -> Result<EagerTensor> {
1062    let mut outputs = outputs.into_iter();
1063    match (outputs.next(), outputs.next()) {
1064        (Some(output), None) => Ok(output),
1065        _ => Err(Error::Internal(format!(
1066            "{name} eager op returned an unexpected number of outputs"
1067        ))),
1068    }
1069}
1070
1071fn two_outputs(outputs: Vec<EagerTensor>, name: &str) -> Result<(EagerTensor, EagerTensor)> {
1072    let mut outputs = outputs.into_iter();
1073    match (outputs.next(), outputs.next(), outputs.next()) {
1074        (Some(lhs), Some(rhs), None) => Ok((lhs, rhs)),
1075        _ => Err(Error::Internal(format!(
1076            "{name} eager op returned an unexpected number of outputs"
1077        ))),
1078    }
1079}