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}