tenferro_linalg/tensor_ext.rs
1//! Receiver-first concrete linear algebra surfaces.
2//!
3//! Owned tensors use [`TensorLinalgExt`], borrowed tensors and views use the
4//! `_read` methods on [`TensorReadLinalgExt`], and typed tensors use
5//! [`TypedTensorLinalgExt`]. All methods dispatch internally to the built-in
6//! CPU/CUDA execution sessions through an erased `&mut dyn BackendSession`
7//! (issue #1680 Phase 3); third-party [`LinalgBackend`] implementations
8//! remain supported through the SPI trait, but the concrete op path is
9//! built-in-session only.
10
11use num_complex::{Complex32, Complex64};
12use tenferro_cpu::with_cpu_exec_session;
13#[cfg(feature = "cuda")]
14use tenferro_gpu::cuda::with_cuda_exec_session;
15use tenferro_tensor::{
16 BackendSession, CompareDir, DType, DotGeneralConfig, Tensor, TensorRead, TensorScalar,
17 TensorWrite, TypedTensor,
18};
19
20use crate::extension::{EighOptions, QrOptions, SvdOptions};
21use crate::{LinalgBackend, RankRevealingQrOptions, RankRevealingQrResult};
22
23/// Scalar types supported by statically typed linear algebra methods.
24///
25/// # Examples
26///
27/// ```rust
28/// use tenferro_linalg::LinalgScalar;
29///
30/// fn accepts_linalg_scalar<T: LinalgScalar>() {}
31/// accepts_linalg_scalar::<f64>();
32/// ```
33pub trait LinalgScalar: TensorScalar + private::Sealed {
34 /// Complex counterpart used by general eigendecomposition.
35 type Complex: TensorScalar;
36}
37
38mod private {
39 pub trait Sealed {}
40 impl Sealed for f32 {}
41 impl Sealed for f64 {}
42 impl Sealed for num_complex::Complex32 {}
43 impl Sealed for num_complex::Complex64 {}
44}
45
46impl LinalgScalar for f32 {
47 type Complex = Complex32;
48}
49impl LinalgScalar for f64 {
50 type Complex = Complex64;
51}
52impl LinalgScalar for Complex32 {
53 type Complex = Complex32;
54}
55impl LinalgScalar for Complex64 {
56 type Complex = Complex64;
57}
58
59/// Fixed typed output tuple for singular value decomposition.
60///
61/// # Examples
62///
63/// ```rust
64/// let _: Option<tenferro_linalg::TypedSvd<f64>> = None;
65/// ```
66pub type TypedSvd<T> = (
67 TypedTensor<T>,
68 TypedTensor<<T as TensorScalar>::Real>,
69 TypedTensor<T>,
70);
71/// Fixed typed output for rank-revealing QR.
72///
73/// # Examples
74///
75/// ```rust
76/// use tenferro_linalg::{RankRevealingQrResult, TypedRankRevealingQrResult};
77/// use tenferro_tensor::TypedTensor;
78///
79/// let result: TypedRankRevealingQrResult<f64> = RankRevealingQrResult {
80/// q: TypedTensor::from_vec_col_major(vec![1, 1], vec![1.0])?,
81/// r: TypedTensor::from_vec_col_major(vec![1, 1], vec![2.0])?,
82/// column_permutation: TypedTensor::from_vec_col_major(vec![1], vec![0_i64])?,
83/// rank: TypedTensor::from_vec_col_major(vec![], vec![1_i64])?,
84/// };
85/// assert_eq!(result.rank.as_slice()?, &[1]);
86/// # Ok::<(), tenferro_tensor::Error>(())
87/// ```
88pub type TypedRankRevealingQrResult<T> = RankRevealingQrResult<TypedTensor<T>, TypedTensor<i64>>;
89/// Fixed typed output tuple for LU decomposition.
90///
91/// # Examples
92///
93/// ```rust
94/// let _: Option<tenferro_linalg::TypedLu<f64>> = None;
95/// ```
96pub type TypedLu<T> = (
97 TypedTensor<T>,
98 TypedTensor<T>,
99 TypedTensor<T>,
100 TypedTensor<<T as TensorScalar>::Real>,
101);
102/// Fixed typed output tuple for complete-pivot LU decomposition.
103///
104/// # Examples
105///
106/// ```rust
107/// let _: Option<tenferro_linalg::TypedFullPivLu<f64>> = None;
108/// ```
109pub type TypedFullPivLu<T> = (
110 TypedTensor<T>,
111 TypedTensor<T>,
112 TypedTensor<T>,
113 TypedTensor<T>,
114 TypedTensor<<T as TensorScalar>::Real>,
115);
116/// Fixed typed output tuple for general eigendecomposition.
117///
118/// # Examples
119///
120/// ```rust
121/// let _: Option<tenferro_linalg::TypedEig<f64>> = None;
122/// ```
123pub type TypedEig<T> = (
124 TypedTensor<<T as LinalgScalar>::Complex>,
125 TypedTensor<<T as LinalgScalar>::Complex>,
126);
127
128/// Linear algebra methods for dtype-erased owned tensors.
129///
130/// # Examples
131///
132/// ```rust
133/// use tenferro_cpu::CpuBackend;
134/// use tenferro_linalg::TensorLinalgExt;
135/// use tenferro_tensor::{BackendSessionHost, Tensor};
136///
137/// let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
138/// let mut host = CpuBackend::new();
139/// let (_u, singular_values, _vt) = host.with_backend_session(|session| a.svd(session))??;
140/// assert_eq!(singular_values.as_slice::<f64>()?, &[4.0, 2.0]);
141/// # Ok::<(), tenferro_tensor::Error>(())
142/// ```
143pub trait TensorLinalgExt {
144 /// # Errors
145 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
146 /// Returns validation, unsupported-backend, numerical, or output-contract errors.
147 /// # Examples
148 ///
149 /// ```rust
150 /// # use tenferro_cpu::CpuBackend;
151 /// # use tenferro_linalg::TensorLinalgExt;
152 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
153 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
154 /// # let mut host = CpuBackend::new();
155 /// let (_u, s, _vt) = host.with_backend_session(|session| a.svd(session))??;
156 /// assert_eq!(s.shape(), &[2]);
157 /// # Ok::<(), tenferro_tensor::Error>(())
158 /// ```
159 fn svd(
160 &self,
161 session: &mut dyn BackendSession,
162 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
163 /// Compute singular values without allocating singular-vector outputs.
164 ///
165 /// # Errors
166 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
167 /// backend has no values-only capability, rather than silently computing a
168 /// full decomposition and discarding its vectors.
169 ///
170 /// # Examples
171 ///
172 /// ```rust
173 /// # use tenferro_cpu::CpuBackend;
174 /// # use tenferro_linalg::TensorLinalgExt;
175 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
176 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
177 /// # let mut host = CpuBackend::new();
178 /// let values = host.with_backend_session(|session| a.svdvals(session))??;
179 /// assert_eq!(values.as_slice::<f64>()?, &[4.0, 2.0]);
180 /// # Ok::<(), tenferro_tensor::Error>(())
181 /// ```
182 fn svdvals(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
183 /// # Errors
184 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
185 /// Returns validation errors for matrix metadata or options, plus SVD backend errors.
186 /// # Examples
187 ///
188 /// ```rust
189 /// # use tenferro_cpu::CpuBackend;
190 /// # use tenferro_linalg::{SvdOptions, TensorLinalgExt};
191 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
192 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
193 /// # let mut host = CpuBackend::new();
194 /// let (_u, s, _vt) = host.with_backend_session(|session| { a.svd_with_options(SvdOptions::default(), session) })??;
195 /// assert_eq!(s.shape(), &[2]);
196 /// # Ok::<(), tenferro_tensor::Error>(())
197 /// ```
198 fn svd_with_options(
199 &self,
200 options: SvdOptions,
201 session: &mut dyn BackendSession,
202 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
203 /// Compute the full-matrices SVD `(U, S, Vt)` with `U` shaped `m x m` and
204 /// `Vt` shaped `n x n`, so the trailing `Vt` rows span the input's right
205 /// nullspace and the trailing `U` columns span its left nullspace.
206 ///
207 /// # Errors
208 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
209 /// backend or CPU provider has no full-matrices kernel; the thin
210 /// decomposition is never substituted for it. Returns validation errors for
211 /// an unsupported rank or dtype, plus backend, numerical, or
212 /// output-contract errors.
213 /// # Examples
214 ///
215 /// ```rust
216 /// # use tenferro_cpu::CpuBackend;
217 /// # use tenferro_linalg::TensorLinalgExt;
218 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
219 /// # let a = Tensor::from_vec_col_major(vec![1, 2], vec![1.0_f64, 1.0])?;
220 /// # let mut host = CpuBackend::new();
221 /// let (u, s, vt) = host.with_backend_session(|session| a.svd_full(session))??;
222 /// assert_eq!(u.shape(), &[1, 1]);
223 /// assert_eq!(s.shape(), &[1]);
224 /// assert_eq!(vt.shape(), &[2, 2]);
225 /// # Ok::<(), tenferro_tensor::Error>(())
226 /// ```
227 fn svd_full(
228 &self,
229 session: &mut dyn BackendSession,
230 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
231 /// # Errors
232 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
233 /// Returns validation, unsupported-backend, numerical, or output-contract errors.
234 /// # Examples
235 ///
236 /// ```rust
237 /// # use tenferro_cpu::CpuBackend;
238 /// # use tenferro_linalg::TensorLinalgExt;
239 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
240 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0])?;
241 /// # let mut host = CpuBackend::new();
242 /// let (q, r) = host.with_backend_session(|session| a.qr(session))??;
243 /// assert_eq!(q.shape(), &[2, 2]);
244 /// assert_eq!(r.shape(), &[2, 2]);
245 /// # Ok::<(), tenferro_tensor::Error>(())
246 /// ```
247 fn qr(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)>;
248
249 /// Initialize opaque compact Householder QR state.
250 ///
251 /// # Errors
252 ///
253 /// Returns validation errors for non-matrix or unsupported-dtype input and
254 /// typed provider errors when compact QR is unavailable.
255 ///
256 /// # Examples
257 ///
258 /// ```rust
259 /// use tenferro_cpu::CpuBackend;
260 /// use tenferro_linalg::TensorLinalgExt;
261 /// use tenferro_tensor::{BackendSessionHost, Tensor};
262 /// let a = Tensor::from_vec_col_major(vec![2, 1], vec![1.0_f64, 2.0])?;
263 /// let mut host = CpuBackend::new();
264 /// let qr = host.with_backend_session(|session| a.householder_qr(session))??;
265 /// assert!(format!("{qr:?}").starts_with("HouseholderQr"));
266 /// # Ok::<(), tenferro_tensor::Error>(())
267 /// ```
268 fn householder_qr(
269 &self,
270 session: &mut dyn BackendSession,
271 ) -> tenferro_tensor::Result<crate::HouseholderQr<Tensor>>;
272
273 /// # Errors
274 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
275 /// Returns validation errors for matrix metadata or options, plus QR backend errors.
276 /// # Examples
277 ///
278 /// ```rust
279 /// # use tenferro_cpu::CpuBackend;
280 /// # use tenferro_linalg::{QrOptions, TensorLinalgExt};
281 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
282 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0])?;
283 /// # let mut host = CpuBackend::new();
284 /// let (q, r) = host.with_backend_session(|session| { a.qr_with_options(QrOptions::default(), session) })??;
285 /// assert_eq!(q.shape(), &[2, 2]);
286 /// assert_eq!(r.shape(), &[2, 2]);
287 /// # Ok::<(), tenferro_tensor::Error>(())
288 /// ```
289 fn qr_with_options(
290 &self,
291 options: QrOptions,
292 session: &mut dyn BackendSession,
293 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
294 /// Compute column-pivoted rank-revealing QR with fixed-shape tensor metadata.
295 ///
296 /// # Errors
297 /// Returns validation errors for rank, dtype, or invalid tolerances;
298 /// numerical failure for non-finite input or diagonal values; and explicit
299 /// unsupported/backend errors when the selected provider cannot execute RRQR.
300 ///
301 /// # Examples
302 ///
303 /// ```rust
304 /// # use tenferro_cpu::CpuBackend;
305 /// # use tenferro_linalg::{RankRevealingQrOptions, TensorLinalgExt};
306 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
307 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
308 /// # let mut host = CpuBackend::new();
309 /// let result = host.with_backend_session(|session| {
310 /// a.rank_revealing_qr(RankRevealingQrOptions::default().rtol(1e-12), session)
311 /// })??;
312 /// assert_eq!(result.column_permutation.shape(), &[2]);
313 /// assert_eq!(result.rank.shape(), &[] as &[usize]);
314 /// # Ok::<(), tenferro_tensor::Error>(())
315 /// ```
316 fn rank_revealing_qr(
317 &self,
318 options: RankRevealingQrOptions,
319 session: &mut dyn BackendSession,
320 ) -> tenferro_tensor::Result<RankRevealingQrResult<Tensor>>;
321 /// # Errors
322 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
323 /// Returns validation, unsupported-backend, numerical, or output-contract errors.
324 /// # Examples
325 ///
326 /// ```rust
327 /// # use tenferro_cpu::CpuBackend;
328 /// # use tenferro_linalg::TensorLinalgExt;
329 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
330 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0])?;
331 /// # let mut host = CpuBackend::new();
332 /// let (_p, l, u, parity) = host.with_backend_session(|session| a.lu(session))??;
333 /// assert_eq!(l.shape(), &[2, 2]);
334 /// assert_eq!(u.shape(), &[2, 2]);
335 /// assert_eq!(parity.shape(), &[] as &[usize]);
336 /// # Ok::<(), tenferro_tensor::Error>(())
337 /// ```
338 fn lu(
339 &self,
340 session: &mut dyn BackendSession,
341 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor)>;
342 /// # Errors
343 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
344 /// Returns validation, unsupported-backend, numerical, or output-contract errors.
345 /// # Examples
346 ///
347 /// ```rust
348 /// # use tenferro_cpu::CpuBackend;
349 /// # use tenferro_linalg::TensorLinalgExt;
350 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
351 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0])?;
352 /// # let mut host = CpuBackend::new();
353 /// let (p, _l, _u, q, parity) = host.with_backend_session(|session| a.full_piv_lu(session))??;
354 /// assert_eq!(p.shape(), &[2, 2]);
355 /// assert_eq!(q.shape(), &[2, 2]);
356 /// assert_eq!(parity.shape(), &[] as &[usize]);
357 /// # Ok::<(), tenferro_tensor::Error>(())
358 /// ```
359 fn full_piv_lu(
360 &self,
361 session: &mut dyn BackendSession,
362 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor, Tensor)>;
363 /// # Errors
364 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
365 /// Returns validation errors for incompatible inputs or a backend/singular-system error.
366 /// # Examples
367 ///
368 /// ```rust
369 /// # use tenferro_cpu::CpuBackend;
370 /// # use tenferro_linalg::TensorLinalgExt;
371 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
372 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![0.0_f64, 2.0, 1.0, 3.0])?;
373 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![-1.0_f64, 5.0])?;
374 /// # let mut host = CpuBackend::new();
375 /// let x = host.with_backend_session(|session| a.full_piv_lu_solve(&b, session))??;
376 /// assert_eq!(x.shape(), &[2, 1]);
377 /// # Ok::<(), tenferro_tensor::Error>(())
378 /// ```
379 fn full_piv_lu_solve(
380 &self,
381 b: &Tensor,
382 session: &mut dyn BackendSession,
383 ) -> tenferro_tensor::Result<Tensor>;
384 /// # Errors
385 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
386 /// Returns validation errors for incompatible inputs or a backend/singular-system error.
387 /// # Examples
388 ///
389 /// ```rust
390 /// # use tenferro_cpu::CpuBackend;
391 /// # use tenferro_linalg::TensorLinalgExt;
392 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
393 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
394 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 8.0])?;
395 /// # let mut host = CpuBackend::new();
396 /// let x = host.with_backend_session(|session| a.solve(&b, session))??;
397 /// assert_eq!(x.as_slice::<f64>()?, &[2.0, 2.0]);
398 /// # Ok::<(), tenferro_tensor::Error>(())
399 /// ```
400 fn solve(
401 &self,
402 b: &Tensor,
403 session: &mut dyn BackendSession,
404 ) -> tenferro_tensor::Result<Tensor>;
405 /// # Errors
406 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
407 /// Returns validation errors for a non-square input or backend/numerical errors.
408 /// # Examples
409 ///
410 /// ```rust
411 /// # use tenferro_cpu::CpuBackend;
412 /// # use tenferro_linalg::TensorLinalgExt;
413 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
414 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![4.0_f64, 2.0, 2.0, 3.0])?;
415 /// # let mut host = CpuBackend::new();
416 /// let l = host.with_backend_session(|session| a.cholesky(session))??;
417 /// assert_eq!(l.shape(), &[2, 2]);
418 /// # Ok::<(), tenferro_tensor::Error>(())
419 /// ```
420 fn cholesky(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
421 /// # Errors
422 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
423 /// Returns validation, unsupported-backend, convergence, or output-contract errors.
424 /// # Examples
425 ///
426 /// ```rust
427 /// # use tenferro_cpu::CpuBackend;
428 /// # use tenferro_linalg::TensorLinalgExt;
429 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
430 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
431 /// # let mut host = CpuBackend::new();
432 /// let (values, vectors) = host.with_backend_session(|session| a.eigh(session))??;
433 /// assert_eq!(values.as_slice::<f64>()?, &[1.0, 3.0]);
434 /// assert_eq!(vectors.shape(), &[2, 2]);
435 /// # Ok::<(), tenferro_tensor::Error>(())
436 /// ```
437 fn eigh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)>;
438 /// # Errors
439 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
440 /// Returns validation errors for matrix metadata or options, plus Eigh backend errors.
441 /// # Examples
442 ///
443 /// ```rust
444 /// # use tenferro_cpu::CpuBackend;
445 /// # use tenferro_linalg::{EighOptions, TensorLinalgExt};
446 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
447 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
448 /// # let mut host = CpuBackend::new();
449 /// let (values, vectors) = host.with_backend_session(|session| { a.eigh_with_options(EighOptions::default(), session) })??;
450 /// assert_eq!(values.shape(), &[2]);
451 /// assert_eq!(vectors.shape(), &[2, 2]);
452 /// # Ok::<(), tenferro_tensor::Error>(())
453 /// ```
454 fn eigh_with_options(
455 &self,
456 options: EighOptions,
457 session: &mut dyn BackendSession,
458 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
459 /// # Errors
460 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
461 /// Returns validation, unsupported-backend, convergence, or output-contract errors.
462 /// # Examples
463 ///
464 /// ```rust
465 /// # use tenferro_cpu::CpuBackend;
466 /// # use tenferro_linalg::TensorLinalgExt;
467 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
468 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
469 /// # let mut host = CpuBackend::new();
470 /// let (values, vectors) = host.with_backend_session(|session| a.eig(session))??;
471 /// assert_eq!(values.shape(), &[2]);
472 /// assert_eq!(vectors.shape(), &[2, 2]);
473 /// # Ok::<(), tenferro_tensor::Error>(())
474 /// ```
475 fn eig(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)>;
476 /// # Errors
477 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
478 /// Returns validation errors for incompatible inputs/flags or backend/numerical errors.
479 #[allow(clippy::too_many_arguments)]
480 /// # Examples
481 ///
482 /// ```rust
483 /// # use tenferro_cpu::CpuBackend;
484 /// # use tenferro_linalg::TensorLinalgExt;
485 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
486 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 1.0, 3.0])?;
487 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 9.0])?;
488 /// # let mut host = CpuBackend::new();
489 /// let x = host.with_backend_session(|session| { a.triangular_solve(&b, true, false, false, false, session) })??;
490 /// assert_eq!(x.shape(), &[2, 1]);
491 /// # Ok::<(), tenferro_tensor::Error>(())
492 /// ```
493 fn triangular_solve(
494 &self,
495 b: &Tensor,
496 left_side: bool,
497 lower: bool,
498 transpose_a: bool,
499 unit_diagonal: bool,
500 session: &mut dyn BackendSession,
501 ) -> tenferro_tensor::Result<Tensor>;
502 /// # Errors
503 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
504 /// Returns validation, unsupported-backend, numerical, or LU output-contract errors.
505 /// # Examples
506 ///
507 /// ```rust
508 /// # use tenferro_cpu::CpuBackend;
509 /// # use tenferro_linalg::TensorLinalgExt;
510 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
511 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
512 /// # let mut host = CpuBackend::new();
513 /// let (sign, logabsdet) = host.with_backend_session(|session| a.slogdet(session))??;
514 /// assert_eq!(sign.as_slice::<f64>()?, &[1.0]);
515 /// assert_eq!(logabsdet.shape(), &[] as &[usize]);
516 /// # Ok::<(), tenferro_tensor::Error>(())
517 /// ```
518 fn slogdet(
519 &self,
520 session: &mut dyn BackendSession,
521 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
522 /// # Errors
523 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
524 /// Returns the validation, backend, numerical, or contract errors from [`Self::slogdet`].
525 /// # Examples
526 ///
527 /// ```rust
528 /// # use tenferro_cpu::CpuBackend;
529 /// # use tenferro_linalg::TensorLinalgExt;
530 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
531 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
532 /// # let mut host = CpuBackend::new();
533 /// let determinant = host.with_backend_session(|session| a.det(session))??;
534 /// assert!((determinant.as_slice::<f64>()?[0] - 8.0).abs() < 1.0e-12);
535 /// # Ok::<(), tenferro_tensor::Error>(())
536 /// ```
537 fn det(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
538 /// # Errors
539 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
540 /// Returns validation, unsupported-backend, or singular-solve numerical errors.
541 /// # Examples
542 ///
543 /// ```rust
544 /// # use tenferro_cpu::CpuBackend;
545 /// # use tenferro_linalg::TensorLinalgExt;
546 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
547 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
548 /// # let mut host = CpuBackend::new();
549 /// let inverse = host.with_backend_session(|session| a.inv(session))??;
550 /// assert_eq!(inverse.as_slice::<f64>()?, &[0.5, 0.0, 0.0, 0.25]);
551 /// # Ok::<(), tenferro_tensor::Error>(())
552 /// ```
553 fn inv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
554 /// # Errors
555 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
556 /// Returns the validation, backend, convergence, or contract errors from [`Self::eigh`].
557 /// # Examples
558 ///
559 /// ```rust
560 /// # use tenferro_cpu::CpuBackend;
561 /// # use tenferro_linalg::TensorLinalgExt;
562 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
563 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
564 /// # let mut host = CpuBackend::new();
565 /// let values = host.with_backend_session(|session| a.eigvalsh(session))??;
566 /// assert_eq!(values.as_slice::<f64>()?, &[1.0, 3.0]);
567 /// # Ok::<(), tenferro_tensor::Error>(())
568 /// ```
569 fn eigvalsh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
570 /// # Errors
571 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
572 /// Returns the validation, backend, convergence, or contract errors from [`Self::eig`].
573 /// # Examples
574 ///
575 /// ```rust
576 /// # use tenferro_cpu::CpuBackend;
577 /// # use tenferro_linalg::TensorLinalgExt;
578 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
579 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
580 /// # let mut host = CpuBackend::new();
581 /// let values = host.with_backend_session(|session| a.eigvals(session))??;
582 /// assert_eq!(values.shape(), &[2]);
583 /// # Ok::<(), tenferro_tensor::Error>(())
584 /// ```
585 fn eigvals(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
586 /// # Errors
587 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
588 /// Returns validation, unsupported-backend, numerical, or SVD contract errors.
589 /// # Examples
590 ///
591 /// ```rust
592 /// # use tenferro_cpu::CpuBackend;
593 /// # use tenferro_linalg::TensorLinalgExt;
594 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
595 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
596 /// # let mut host = CpuBackend::new();
597 /// let pseudoinverse = host.with_backend_session(|session| a.pinv(session))??;
598 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
599 /// # Ok::<(), tenferro_tensor::Error>(())
600 /// ```
601 fn pinv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
602 /// # Errors
603 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
604 /// Returns a validation error for invalid `rtol`, plus errors from [`Self::pinv`].
605 /// # Examples
606 ///
607 /// ```rust
608 /// # use tenferro_cpu::CpuBackend;
609 /// # use tenferro_linalg::TensorLinalgExt;
610 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
611 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
612 /// # let mut host = CpuBackend::new();
613 /// let pseudoinverse = host.with_backend_session(|session| a.pinv_with_rtol(1.0e-12, session))??;
614 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
615 /// # Ok::<(), tenferro_tensor::Error>(())
616 /// ```
617 fn pinv_with_rtol(
618 &self,
619 rtol: f64,
620 session: &mut dyn BackendSession,
621 ) -> tenferro_tensor::Result<Tensor>;
622 /// # Errors
623 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
624 /// Returns validation errors for axes/order combinations or required backend operations.
625 /// # Examples
626 ///
627 /// ```rust
628 /// # use tenferro_cpu::CpuBackend;
629 /// # use tenferro_linalg::TensorLinalgExt;
630 /// # use tenferro_tensor::{BackendSessionHost, Tensor};
631 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![3.0_f64, 0.0, 0.0, 4.0])?;
632 /// # let mut host = CpuBackend::new();
633 /// let frobenius = host.with_backend_session(|session| a.norm(None, None, false, session))??;
634 /// assert_eq!(frobenius.shape(), &[] as &[usize]);
635 /// # Ok::<(), tenferro_tensor::Error>(())
636 /// ```
637 fn norm(
638 &self,
639 ord: Option<f64>,
640 dim: Option<&[usize]>,
641 keepdim: bool,
642 session: &mut dyn BackendSession,
643 ) -> tenferro_tensor::Result<Tensor>;
644}
645
646/// Linear algebra methods for borrowed tensor reads.
647///
648/// # Examples
649///
650/// ```rust
651/// use tenferro_cpu::CpuBackend;
652/// use tenferro_linalg::TensorReadLinalgExt;
653/// use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
654///
655/// let input = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
656/// let mut host = CpuBackend::new();
657/// let (_q, r) = host.with_backend_session(|session| {
658/// TensorRead::from_tensor(&input).qr_read(session)
659/// })??;
660/// assert_eq!(r.shape(), &[2, 2]);
661/// # Ok::<(), tenferro_tensor::Error>(())
662/// ```
663pub trait TensorReadLinalgExt {
664 /// # Errors
665 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
666 /// Returns validation, same-placement materialization, backend, numerical, or contract errors.
667 /// # Examples
668 ///
669 /// ```rust
670 /// # use tenferro_cpu::CpuBackend;
671 /// # use tenferro_linalg::TensorReadLinalgExt;
672 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
673 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
674 /// # let mut host = CpuBackend::new();
675 /// let (_u, s, _vt) = host.with_backend_session(|session| {
676 /// TensorRead::from_tensor(&a).svd_read(session)
677 /// })??;
678 /// assert_eq!(s.shape(), &[2]);
679 /// # Ok::<(), tenferro_tensor::Error>(())
680 /// ```
681 fn svd_read(
682 self,
683 session: &mut dyn BackendSession,
684 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
685 /// Compute singular values from a borrowed tensor read target.
686 ///
687 /// Eligible faer host views are consumed without a full input copy. A
688 /// provider that requires owned compact storage may materialize explicitly;
689 /// unsupported providers return a typed error.
690 ///
691 /// # Errors
692 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
693 /// backend has no borrowed values-only capability.
694 ///
695 /// # Examples
696 ///
697 /// ```rust
698 /// # use tenferro_cpu::CpuBackend;
699 /// # use tenferro_linalg::TensorReadLinalgExt;
700 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
701 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
702 /// # let mut host = CpuBackend::new();
703 /// let values = host.with_backend_session(|session| TensorRead::from_tensor(&a).svdvals_read(session))??;
704 /// assert_eq!(values.as_slice::<f64>()?, &[4.0, 2.0]);
705 /// # Ok::<(), tenferro_tensor::Error>(())
706 /// ```
707 fn svdvals_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
708 /// # Errors
709 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
710 /// Returns validation errors for metadata/options, plus read/materialization or SVD errors.
711 /// # Examples
712 ///
713 /// ```rust
714 /// # use tenferro_cpu::CpuBackend;
715 /// # use tenferro_linalg::{SvdOptions, TensorReadLinalgExt};
716 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
717 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
718 /// # let mut host = CpuBackend::new();
719 /// let (_u, s, _vt) = host.with_backend_session(|session| {
720 /// TensorRead::from_tensor(&a).svd_with_options_read(SvdOptions::default(), session)
721 /// })??;
722 /// assert_eq!(s.shape(), &[2]);
723 /// # Ok::<(), tenferro_tensor::Error>(())
724 /// ```
725 fn svd_with_options_read(
726 self,
727 options: SvdOptions,
728 session: &mut dyn BackendSession,
729 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
730 /// Compute the full-matrices SVD `(U, S, Vt)` from a borrowed tensor read
731 /// target, with `U` shaped `m x m` and `Vt` shaped `n x n`.
732 ///
733 /// Eligible faer host views (rank 2, host placement, non-negative strides)
734 /// are consumed without an input copy. A provider that requires owned
735 /// compact storage materializes explicitly inside the provider boundary;
736 /// unsupported providers return a typed error instead of a thin result. The
737 /// borrowed source is never modified.
738 ///
739 /// # Errors
740 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
741 /// backend or CPU provider has no borrowed full-matrices capability.
742 /// Returns validation errors for an unsupported rank, dtype, or placement,
743 /// plus same-placement materialization, backend, numerical, or
744 /// output-contract errors.
745 /// # Examples
746 ///
747 /// ```rust
748 /// # use tenferro_cpu::CpuBackend;
749 /// # use tenferro_linalg::TensorReadLinalgExt;
750 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
751 /// # let a = Tensor::from_vec_col_major(vec![1, 2], vec![1.0_f64, 1.0])?;
752 /// # let mut host = CpuBackend::new();
753 /// let (u, s, vt) = host.with_backend_session(|session| {
754 /// TensorRead::from_tensor(&a).svd_full_read(session)
755 /// })??;
756 /// assert_eq!(u.shape(), &[1, 1]);
757 /// assert_eq!(s.shape(), &[1]);
758 /// assert_eq!(vt.shape(), &[2, 2]);
759 /// # Ok::<(), tenferro_tensor::Error>(())
760 /// ```
761 fn svd_full_read(
762 self,
763 session: &mut dyn BackendSession,
764 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)>;
765 /// # Errors
766 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
767 /// Returns validation, same-placement materialization, backend, numerical, or contract errors.
768 /// # Examples
769 ///
770 /// ```rust
771 /// # use tenferro_cpu::CpuBackend;
772 /// # use tenferro_linalg::TensorReadLinalgExt;
773 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
774 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0])?;
775 /// # let mut host = CpuBackend::new();
776 /// let (q, r) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).qr_read(session) })??;
777 /// assert_eq!(q.shape(), &[2, 2]);
778 /// assert_eq!(r.shape(), &[2, 2]);
779 /// # Ok::<(), tenferro_tensor::Error>(())
780 /// ```
781 fn qr_read(self, session: &mut dyn BackendSession)
782 -> tenferro_tensor::Result<(Tensor, Tensor)>;
783 /// # Errors
784 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
785 /// Returns validation errors for metadata/options, plus read/materialization or QR errors.
786 /// # Examples
787 ///
788 /// ```rust
789 /// # use tenferro_cpu::CpuBackend;
790 /// # use tenferro_linalg::{QrOptions, TensorReadLinalgExt};
791 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
792 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 1.0])?;
793 /// # let mut host = CpuBackend::new();
794 /// let (q, r) = host.with_backend_session(|session| {
795 /// TensorRead::from_tensor(&a).qr_with_options_read(QrOptions::default(), session)
796 /// })??;
797 /// assert_eq!(q.shape(), &[2, 2]);
798 /// assert_eq!(r.shape(), &[2, 2]);
799 /// # Ok::<(), tenferro_tensor::Error>(())
800 /// ```
801 fn qr_with_options_read(
802 self,
803 options: QrOptions,
804 session: &mut dyn BackendSession,
805 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
806 /// Compute rank-revealing QR from a borrowed tensor read.
807 ///
808 /// # Errors
809 /// Returns validation, same-placement materialization, numerical, backend,
810 /// or explicit unsupported errors.
811 ///
812 /// # Examples
813 ///
814 /// ```rust
815 /// # use tenferro_cpu::CpuBackend;
816 /// # use tenferro_linalg::{RankRevealingQrOptions, TensorReadLinalgExt};
817 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
818 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
819 /// # let mut host = CpuBackend::new();
820 /// let result = host.with_backend_session(|session| {
821 /// TensorRead::from_tensor(&a).rank_revealing_qr_read(
822 /// RankRevealingQrOptions::default(), session)
823 /// })??;
824 /// assert_eq!(result.rank.as_slice::<i64>()?, &[2]);
825 /// # Ok::<(), tenferro_tensor::Error>(())
826 /// ```
827 fn rank_revealing_qr_read(
828 self,
829 options: RankRevealingQrOptions,
830 session: &mut dyn BackendSession,
831 ) -> tenferro_tensor::Result<RankRevealingQrResult<Tensor>>;
832 /// # Errors
833 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
834 /// Returns validation, same-placement materialization, backend, numerical, or contract errors.
835 /// # Examples
836 ///
837 /// ```rust
838 /// # use tenferro_cpu::CpuBackend;
839 /// # use tenferro_linalg::TensorReadLinalgExt;
840 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
841 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0])?;
842 /// # let mut host = CpuBackend::new();
843 /// let (_p, l, u, _parity) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).lu_read(session) })??;
844 /// assert_eq!(l.shape(), &[2, 2]);
845 /// assert_eq!(u.shape(), &[2, 2]);
846 /// # Ok::<(), tenferro_tensor::Error>(())
847 /// ```
848 fn lu_read(
849 self,
850 session: &mut dyn BackendSession,
851 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor)>;
852 /// # Errors
853 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
854 /// Returns validation, same-placement materialization, backend, numerical, or contract errors.
855 /// # Examples
856 ///
857 /// ```rust
858 /// # use tenferro_cpu::CpuBackend;
859 /// # use tenferro_linalg::TensorReadLinalgExt;
860 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
861 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 3.0, 2.0, 4.0])?;
862 /// # let mut host = CpuBackend::new();
863 /// let (p, _l, _u, q, _parity) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).full_piv_lu_read(session) })??;
864 /// assert_eq!(p.shape(), &[2, 2]);
865 /// assert_eq!(q.shape(), &[2, 2]);
866 /// # Ok::<(), tenferro_tensor::Error>(())
867 /// ```
868 fn full_piv_lu_read(
869 self,
870 session: &mut dyn BackendSession,
871 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor, Tensor)>;
872 /// # Errors
873 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
874 /// Returns incompatible-input validation, read/materialization, backend, or singular errors.
875 /// # Examples
876 ///
877 /// ```rust
878 /// # use tenferro_cpu::CpuBackend;
879 /// # use tenferro_linalg::TensorReadLinalgExt;
880 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
881 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![0.0_f64, 2.0, 1.0, 3.0])?;
882 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![-1.0_f64, 5.0])?;
883 /// # let mut host = CpuBackend::new();
884 /// let x = host.with_backend_session(|session| {
885 /// TensorRead::from_tensor(&a).full_piv_lu_solve_read(TensorRead::from_tensor(&b), session)
886 /// })??;
887 /// assert_eq!(x.shape(), &[2, 1]);
888 /// # Ok::<(), tenferro_tensor::Error>(())
889 /// ```
890 fn full_piv_lu_solve_read(
891 self,
892 b: TensorRead<'_>,
893 session: &mut dyn BackendSession,
894 ) -> tenferro_tensor::Result<Tensor>;
895 /// # Errors
896 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
897 /// Returns incompatible-input validation, read/materialization, backend, or singular errors.
898 /// # Examples
899 ///
900 /// ```rust
901 /// # use tenferro_cpu::CpuBackend;
902 /// # use tenferro_linalg::TensorReadLinalgExt;
903 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
904 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
905 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 8.0])?;
906 /// # let mut host = CpuBackend::new();
907 /// let x = host.with_backend_session(|session| { TensorRead::from_tensor(&a).solve_read(TensorRead::from_tensor(&b), session) })??;
908 /// assert_eq!(x.as_slice::<f64>()?, &[2.0, 2.0]);
909 /// # Ok::<(), tenferro_tensor::Error>(())
910 /// ```
911 fn solve_read(
912 self,
913 b: TensorRead<'_>,
914 session: &mut dyn BackendSession,
915 ) -> tenferro_tensor::Result<Tensor>;
916 /// Solve into a caller-owned destination without allocating the result at
917 /// the public API boundary.
918 ///
919 /// Backends with a native path may write directly into a compatible target;
920 /// the trait default preserves the ordinary solve-read plus copy behavior.
921 ///
922 /// # Errors
923 /// Returns `tenferro_tensor_core::ShapeMismatch` or
924 /// `tenferro_tensor_core::ValidationError::DTypeMismatch` for incompatible
925 /// destination metadata, `tenferro_tensor_core::ValidationError::InvalidArgument`
926 /// for aliasing or placement violations, `Error::Unsupported` when the
927 /// extension implementation or provider is unavailable, and `Error::Singular`
928 /// for a singular system.
929 /// # Examples
930 ///
931 /// ```rust
932 /// # use tenferro_cpu::CpuBackend;
933 /// # use tenferro_linalg::TensorReadLinalgExt;
934 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead, TensorWrite};
935 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
936 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 8.0])?;
937 /// # let mut out = Tensor::from_vec_col_major(vec![2, 1], vec![0.0_f64; 2])?;
938 /// # let mut host = CpuBackend::new();
939 /// host.with_backend_session(|session| {
940 /// TensorRead::from_tensor(&a).solve_read_into(
941 /// TensorRead::from_tensor(&b),
942 /// TensorWrite::from_tensor(&mut out),
943 /// session,
944 /// )
945 /// })??;
946 /// assert_eq!(out.as_slice::<f64>()?, &[2.0, 2.0]);
947 /// # Ok::<(), tenferro_tensor::Error>(())
948 /// ```
949 fn solve_read_into(
950 self,
951 b: TensorRead<'_>,
952 out: TensorWrite<'_>,
953 session: &mut dyn BackendSession,
954 ) -> tenferro_tensor::Result<()>
955 where
956 Self: Sized,
957 {
958 let _ = (self, b, out, session);
959 Err(tenferro_tensor::Error::unsupported(
960 "solve_read_into",
961 "this tensor-read extension implementation does not accept borrowed solve targets",
962 ))
963 }
964 /// # Errors
965 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
966 /// Returns matrix validation, read/materialization, backend, or positive-definiteness errors.
967 /// # Examples
968 ///
969 /// ```rust
970 /// # use tenferro_cpu::CpuBackend;
971 /// # use tenferro_linalg::TensorReadLinalgExt;
972 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
973 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![4.0_f64, 2.0, 2.0, 3.0])?;
974 /// # let mut host = CpuBackend::new();
975 /// let l = host.with_backend_session(|session| { TensorRead::from_tensor(&a).cholesky_read(session) })??;
976 /// assert_eq!(l.shape(), &[2, 2]);
977 /// # Ok::<(), tenferro_tensor::Error>(())
978 /// ```
979 fn cholesky_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
980 /// # Errors
981 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
982 /// Returns validation, read/materialization, backend, convergence, or contract errors.
983 /// # Examples
984 ///
985 /// ```rust
986 /// # use tenferro_cpu::CpuBackend;
987 /// # use tenferro_linalg::TensorReadLinalgExt;
988 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
989 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
990 /// # let mut host = CpuBackend::new();
991 /// let (values, vectors) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).eigh_read(session) })??;
992 /// assert_eq!(values.as_slice::<f64>()?, &[1.0, 3.0]);
993 /// assert_eq!(vectors.shape(), &[2, 2]);
994 /// # Ok::<(), tenferro_tensor::Error>(())
995 /// ```
996 fn eigh_read(
997 self,
998 session: &mut dyn BackendSession,
999 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
1000 /// # Errors
1001 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1002 /// Returns validation errors for metadata/options, plus read or Eigh backend errors.
1003 /// # Examples
1004 ///
1005 /// ```rust
1006 /// # use tenferro_cpu::CpuBackend;
1007 /// # use tenferro_linalg::{EighOptions, TensorReadLinalgExt};
1008 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1009 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
1010 /// # let mut host = CpuBackend::new();
1011 /// let (values, vectors) = host.with_backend_session(|session| {
1012 /// TensorRead::from_tensor(&a).eigh_with_options_read(EighOptions::default(), session)
1013 /// })??;
1014 /// assert_eq!(values.shape(), &[2]);
1015 /// assert_eq!(vectors.shape(), &[2, 2]);
1016 /// # Ok::<(), tenferro_tensor::Error>(())
1017 /// ```
1018 fn eigh_with_options_read(
1019 self,
1020 options: EighOptions,
1021 session: &mut dyn BackendSession,
1022 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
1023 /// # Errors
1024 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1025 /// Returns validation, read/materialization, backend, convergence, or contract errors.
1026 /// # Examples
1027 ///
1028 /// ```rust
1029 /// # use tenferro_cpu::CpuBackend;
1030 /// # use tenferro_linalg::TensorReadLinalgExt;
1031 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1032 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
1033 /// # let mut host = CpuBackend::new();
1034 /// let (values, vectors) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).eig_read(session) })??;
1035 /// assert_eq!(values.shape(), &[2]);
1036 /// assert_eq!(vectors.shape(), &[2, 2]);
1037 /// # Ok::<(), tenferro_tensor::Error>(())
1038 /// ```
1039 fn eig_read(
1040 self,
1041 session: &mut dyn BackendSession,
1042 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
1043 /// # Errors
1044 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1045 /// Returns incompatible-input/flag validation, read, backend, or singular errors.
1046 #[allow(clippy::too_many_arguments)]
1047 /// # Examples
1048 ///
1049 /// ```rust
1050 /// # use tenferro_cpu::CpuBackend;
1051 /// # use tenferro_linalg::TensorReadLinalgExt;
1052 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1053 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 1.0, 3.0])?;
1054 /// # let b = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 9.0])?;
1055 /// # let mut host = CpuBackend::new();
1056 /// let x = host.with_backend_session(|session| {
1057 /// TensorRead::from_tensor(&a).triangular_solve_read(
1058 /// TensorRead::from_tensor(&b),
1059 /// true,
1060 /// false,
1061 /// false,
1062 /// false,
1063 /// session,
1064 /// )
1065 /// })??;
1066 /// assert_eq!(x.shape(), &[2, 1]);
1067 /// # Ok::<(), tenferro_tensor::Error>(())
1068 /// ```
1069 fn triangular_solve_read(
1070 self,
1071 b: TensorRead<'_>,
1072 left_side: bool,
1073 lower: bool,
1074 transpose_a: bool,
1075 unit_diagonal: bool,
1076 session: &mut dyn BackendSession,
1077 ) -> tenferro_tensor::Result<Tensor>;
1078 /// # Errors
1079 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1080 /// Returns validation, read/materialization, backend, numerical, or LU contract errors.
1081 /// # Examples
1082 ///
1083 /// ```rust
1084 /// # use tenferro_cpu::CpuBackend;
1085 /// # use tenferro_linalg::TensorReadLinalgExt;
1086 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1087 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
1088 /// # let mut host = CpuBackend::new();
1089 /// let (sign, logabsdet) = host.with_backend_session(|session| { TensorRead::from_tensor(&a).slogdet_read(session) })??;
1090 /// assert_eq!(sign.as_slice::<f64>()?, &[1.0]);
1091 /// assert_eq!(logabsdet.shape(), &[] as &[usize]);
1092 /// # Ok::<(), tenferro_tensor::Error>(())
1093 /// ```
1094 fn slogdet_read(
1095 self,
1096 session: &mut dyn BackendSession,
1097 ) -> tenferro_tensor::Result<(Tensor, Tensor)>;
1098 /// # Errors
1099 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1100 /// Returns validation, read/materialization, backend, numerical, or contract errors.
1101 /// # Examples
1102 ///
1103 /// ```rust
1104 /// # use tenferro_cpu::CpuBackend;
1105 /// # use tenferro_linalg::TensorReadLinalgExt;
1106 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1107 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
1108 /// # let mut host = CpuBackend::new();
1109 /// let determinant = host.with_backend_session(|session| { TensorRead::from_tensor(&a).det_read(session) })??;
1110 /// assert!((determinant.as_slice::<f64>()?[0] - 8.0).abs() < 1.0e-12);
1111 /// # Ok::<(), tenferro_tensor::Error>(())
1112 /// ```
1113 fn det_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1114 /// # Errors
1115 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1116 /// Returns validation, read/materialization, backend, or singular-solve errors.
1117 /// # Examples
1118 ///
1119 /// ```rust
1120 /// # use tenferro_cpu::CpuBackend;
1121 /// # use tenferro_linalg::TensorReadLinalgExt;
1122 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1123 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
1124 /// # let mut host = CpuBackend::new();
1125 /// let inverse = host.with_backend_session(|session| { TensorRead::from_tensor(&a).inv_read(session) })??;
1126 /// assert_eq!(inverse.as_slice::<f64>()?, &[0.5, 0.0, 0.0, 0.25]);
1127 /// # Ok::<(), tenferro_tensor::Error>(())
1128 /// ```
1129 fn inv_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1130 /// # Errors
1131 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1132 /// Returns validation, read/materialization, backend, convergence, or contract errors.
1133 /// # Examples
1134 ///
1135 /// ```rust
1136 /// # use tenferro_cpu::CpuBackend;
1137 /// # use tenferro_linalg::TensorReadLinalgExt;
1138 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1139 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 3.0])?;
1140 /// # let mut host = CpuBackend::new();
1141 /// let values = host.with_backend_session(|session| { TensorRead::from_tensor(&a).eigvalsh_read(session) })??;
1142 /// assert_eq!(values.as_slice::<f64>()?, &[1.0, 3.0]);
1143 /// # Ok::<(), tenferro_tensor::Error>(())
1144 /// ```
1145 fn eigvalsh_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1146 /// # Errors
1147 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1148 /// Returns validation, read/materialization, backend, convergence, or contract errors.
1149 /// # Examples
1150 ///
1151 /// ```rust
1152 /// # use tenferro_cpu::CpuBackend;
1153 /// # use tenferro_linalg::TensorReadLinalgExt;
1154 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1155 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 0.0, 0.0, 2.0])?;
1156 /// # let mut host = CpuBackend::new();
1157 /// let values = host.with_backend_session(|session| { TensorRead::from_tensor(&a).eigvals_read(session) })??;
1158 /// assert_eq!(values.shape(), &[2]);
1159 /// # Ok::<(), tenferro_tensor::Error>(())
1160 /// ```
1161 fn eigvals_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1162 /// # Errors
1163 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1164 /// Returns validation, read/materialization, backend, numerical, or SVD contract errors.
1165 /// # Examples
1166 ///
1167 /// ```rust
1168 /// # use tenferro_cpu::CpuBackend;
1169 /// # use tenferro_linalg::TensorReadLinalgExt;
1170 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1171 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
1172 /// # let mut host = CpuBackend::new();
1173 /// let pseudoinverse = host.with_backend_session(|session| { TensorRead::from_tensor(&a).pinv_read(session) })??;
1174 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
1175 /// # Ok::<(), tenferro_tensor::Error>(())
1176 /// ```
1177 fn pinv_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1178 /// # Errors
1179 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1180 /// Returns a validation error for invalid `rtol`, plus errors from [`Self::pinv_read`].
1181 /// # Examples
1182 ///
1183 /// ```rust
1184 /// # use tenferro_cpu::CpuBackend;
1185 /// # use tenferro_linalg::TensorReadLinalgExt;
1186 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1187 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![2.0_f64, 0.0, 0.0, 4.0])?;
1188 /// # let mut host = CpuBackend::new();
1189 /// let pseudoinverse = host.with_backend_session(|session| { TensorRead::from_tensor(&a).pinv_with_rtol_read(1.0e-12, session) })??;
1190 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
1191 /// # Ok::<(), tenferro_tensor::Error>(())
1192 /// ```
1193 fn pinv_with_rtol_read(
1194 self,
1195 rtol: f64,
1196 session: &mut dyn BackendSession,
1197 ) -> tenferro_tensor::Result<Tensor>;
1198 /// # Errors
1199 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1200 /// Returns validation errors for axes/order combinations or required read/backend operations.
1201 /// # Examples
1202 ///
1203 /// ```rust
1204 /// # use tenferro_cpu::CpuBackend;
1205 /// # use tenferro_linalg::TensorReadLinalgExt;
1206 /// # use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1207 /// # let a = Tensor::from_vec_col_major(vec![2, 2], vec![3.0_f64, 0.0, 0.0, 4.0])?;
1208 /// # let mut host = CpuBackend::new();
1209 /// let frobenius = host.with_backend_session(|session| { TensorRead::from_tensor(&a).norm_read(None, None, false, session) })??;
1210 /// assert_eq!(frobenius.shape(), &[] as &[usize]);
1211 /// # Ok::<(), tenferro_tensor::Error>(())
1212 /// ```
1213 fn norm_read(
1214 self,
1215 ord: Option<f64>,
1216 dim: Option<&[usize]>,
1217 keepdim: bool,
1218 session: &mut dyn BackendSession,
1219 ) -> tenferro_tensor::Result<Tensor>;
1220}
1221
1222/// Linear algebra methods with statically typed inputs and outputs.
1223///
1224/// Singular values, Hermitian eigenvalues, determinant log-magnitudes, and
1225/// norms use `T::Real`; general eigenvalues and eigenvectors use `T::Complex`.
1226///
1227/// # Examples
1228///
1229/// ```rust
1230/// use tenferro_cpu::CpuBackend;
1231/// use tenferro_linalg::TypedTensorLinalgExt;
1232/// use tenferro_tensor::{BackendSessionHost, TypedTensor};
1233///
1234/// let input = TypedTensor::<f64>::from_vec_col_major(
1235/// vec![2, 2],
1236/// vec![2.0, 0.0, 0.0, 4.0],
1237/// )?;
1238/// let mut host = CpuBackend::new();
1239/// let (_u, singular_values, _vt) = host.with_backend_session(|session| input.svd(session))??;
1240/// assert_eq!(singular_values.as_slice()?, &[4.0, 2.0]);
1241/// # Ok::<(), tenferro_tensor::Error>(())
1242/// ```
1243pub trait TypedTensorLinalgExt<T: LinalgScalar> {
1244 /// # Errors
1245 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1246 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1247 /// # Examples
1248 ///
1249 /// ```rust
1250 /// # use tenferro_cpu::CpuBackend;
1251 /// # use tenferro_linalg::TypedTensorLinalgExt;
1252 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1253 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1254 /// # let mut host = CpuBackend::new();
1255 /// let (_u, s, _vt) = host.with_backend_session(|session| a.svd(session))??;
1256 /// assert_eq!(s.as_slice()?, &[4.0, 2.0]);
1257 /// # Ok::<(), tenferro_tensor::Error>(())
1258 /// ```
1259 fn svd(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedSvd<T>>;
1260 /// Compute singular values without allocating singular-vector outputs.
1261 ///
1262 /// # Errors
1263 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
1264 /// backend has no values-only capability.
1265 ///
1266 /// # Examples
1267 ///
1268 /// ```rust
1269 /// # use tenferro_cpu::CpuBackend;
1270 /// # use tenferro_linalg::TypedTensorLinalgExt;
1271 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1272 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1273 /// # let mut host = CpuBackend::new();
1274 /// let values = host.with_backend_session(|session| a.svdvals(session))??;
1275 /// assert_eq!(values.as_slice()?, &[4.0, 2.0]);
1276 /// # Ok::<(), tenferro_tensor::Error>(())
1277 /// ```
1278 fn svdvals(
1279 &self,
1280 session: &mut dyn BackendSession,
1281 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>>;
1282 /// # Errors
1283 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1284 /// Returns validation errors for metadata/options, plus backend or typed-output errors.
1285 /// # Examples
1286 ///
1287 /// ```rust
1288 /// # use tenferro_cpu::CpuBackend;
1289 /// # use tenferro_linalg::{SvdOptions, TypedTensorLinalgExt};
1290 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1291 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1292 /// # let mut host = CpuBackend::new();
1293 /// let (_u, s, _vt) = host.with_backend_session(|session| { a.svd_with_options(SvdOptions::default(), session) })??;
1294 /// assert_eq!(s.as_slice()?, &[4.0, 2.0]);
1295 /// # Ok::<(), tenferro_tensor::Error>(())
1296 /// ```
1297 fn svd_with_options(
1298 &self,
1299 options: SvdOptions,
1300 session: &mut dyn BackendSession,
1301 ) -> tenferro_tensor::Result<TypedSvd<T>>;
1302 /// Compute the full-matrices SVD `(U, S, Vt)` with `U` shaped `m x m` and
1303 /// `Vt` shaped `n x n`. The singular values keep the real counterpart dtype
1304 /// of `T`, exactly as [`TypedTensorLinalgExt::svd`] does.
1305 ///
1306 /// # Errors
1307 /// Returns [`tenferro_tensor::Error::Unsupported`] when the selected
1308 /// backend or CPU provider has no full-matrices kernel. Returns validation
1309 /// errors for an unsupported rank or dtype, plus backend, numerical,
1310 /// output-contract, or typed-downcast errors.
1311 /// # Examples
1312 ///
1313 /// ```rust
1314 /// # use tenferro_cpu::CpuBackend;
1315 /// # use tenferro_linalg::TypedTensorLinalgExt;
1316 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1317 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![1, 2], vec![1.0, 1.0])?;
1318 /// # let mut host = CpuBackend::new();
1319 /// let (u, s, vt) = host.with_backend_session(|session| a.svd_full(session))??;
1320 /// assert_eq!(u.shape(), &[1, 1]);
1321 /// assert_eq!(vt.shape(), &[2, 2]);
1322 /// assert!((s.as_slice()?[0] - 2.0_f64.sqrt()).abs() < 1e-12);
1323 /// # Ok::<(), tenferro_tensor::Error>(())
1324 /// ```
1325 fn svd_full(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedSvd<T>>;
1326 /// # Errors
1327 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1328 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1329 /// # Examples
1330 ///
1331 /// ```rust
1332 /// # use tenferro_cpu::CpuBackend;
1333 /// # use tenferro_linalg::TypedTensorLinalgExt;
1334 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1335 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 1.0])?;
1336 /// # let mut host = CpuBackend::new();
1337 /// let (q, r) = host.with_backend_session(|session| a.qr(session))??;
1338 /// assert_eq!(q.shape(), &[2, 2]);
1339 /// assert_eq!(r.shape(), &[2, 2]);
1340 /// # Ok::<(), tenferro_tensor::Error>(())
1341 /// ```
1342 fn qr(
1343 &self,
1344 session: &mut dyn BackendSession,
1345 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<T>)>;
1346 /// # Errors
1347 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1348 /// Returns validation errors for metadata/options, plus backend or typed-output errors.
1349 /// # Examples
1350 ///
1351 /// ```rust
1352 /// # use tenferro_cpu::CpuBackend;
1353 /// # use tenferro_linalg::{QrOptions, TypedTensorLinalgExt};
1354 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1355 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 1.0])?;
1356 /// # let mut host = CpuBackend::new();
1357 /// let (q, r) = host.with_backend_session(|session| { a.qr_with_options(QrOptions::default(), session) })??;
1358 /// assert_eq!(q.shape(), &[2, 2]);
1359 /// assert_eq!(r.shape(), &[2, 2]);
1360 /// # Ok::<(), tenferro_tensor::Error>(())
1361 /// ```
1362 fn qr_with_options(
1363 &self,
1364 options: QrOptions,
1365 session: &mut dyn BackendSession,
1366 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<T>)>;
1367 /// Compute typed rank-revealing QR with `i64` permutation and rank tensors.
1368 ///
1369 /// # Errors
1370 /// Returns validation, numerical, backend, unsupported, or typed-output
1371 /// contract errors.
1372 ///
1373 /// # Examples
1374 ///
1375 /// ```rust
1376 /// # use tenferro_cpu::CpuBackend;
1377 /// # use tenferro_linalg::{RankRevealingQrOptions, TypedTensorLinalgExt};
1378 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1379 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 2.0])?;
1380 /// # let mut host = CpuBackend::new();
1381 /// let result = host.with_backend_session(|session| {
1382 /// a.rank_revealing_qr(RankRevealingQrOptions::default(), session)
1383 /// })??;
1384 /// assert_eq!(result.rank.as_slice()?, &[2_i64]);
1385 /// # Ok::<(), tenferro_tensor::Error>(())
1386 /// ```
1387 fn rank_revealing_qr(
1388 &self,
1389 options: RankRevealingQrOptions,
1390 session: &mut dyn BackendSession,
1391 ) -> tenferro_tensor::Result<TypedRankRevealingQrResult<T>>;
1392 /// # Errors
1393 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1394 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1395 /// # Examples
1396 ///
1397 /// ```rust
1398 /// # use tenferro_cpu::CpuBackend;
1399 /// # use tenferro_linalg::TypedTensorLinalgExt;
1400 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1401 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 3.0, 2.0, 4.0])?;
1402 /// # let mut host = CpuBackend::new();
1403 /// let (_p, l, u, _parity) = host.with_backend_session(|session| a.lu(session))??;
1404 /// assert_eq!(l.shape(), &[2, 2]);
1405 /// assert_eq!(u.shape(), &[2, 2]);
1406 /// # Ok::<(), tenferro_tensor::Error>(())
1407 /// ```
1408 fn lu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedLu<T>>;
1409 /// # Errors
1410 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1411 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1412 /// # Examples
1413 ///
1414 /// ```rust
1415 /// # use tenferro_cpu::CpuBackend;
1416 /// # use tenferro_linalg::TypedTensorLinalgExt;
1417 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1418 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 3.0, 2.0, 4.0])?;
1419 /// # let mut host = CpuBackend::new();
1420 /// let (p, _l, _u, q, _parity) = host.with_backend_session(|session| a.full_piv_lu(session))??;
1421 /// assert_eq!(p.shape(), &[2, 2]);
1422 /// assert_eq!(q.shape(), &[2, 2]);
1423 /// # Ok::<(), tenferro_tensor::Error>(())
1424 /// ```
1425 fn full_piv_lu(
1426 &self,
1427 session: &mut dyn BackendSession,
1428 ) -> tenferro_tensor::Result<TypedFullPivLu<T>>;
1429 /// # Errors
1430 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1431 /// Returns incompatible-input validation, backend, singular, or typed-downcast errors.
1432 /// # Examples
1433 ///
1434 /// ```rust
1435 /// # use tenferro_cpu::CpuBackend;
1436 /// # use tenferro_linalg::TypedTensorLinalgExt;
1437 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1438 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![0.0, 2.0, 1.0, 3.0])?;
1439 /// # let b = TypedTensor::<f64>::from_vec_col_major(vec![2, 1], vec![-1.0, 5.0])?;
1440 /// # let mut host = CpuBackend::new();
1441 /// let x = host.with_backend_session(|session| a.full_piv_lu_solve(&b, session))??;
1442 /// assert_eq!(x.shape(), &[2, 1]);
1443 /// # Ok::<(), tenferro_tensor::Error>(())
1444 /// ```
1445 fn full_piv_lu_solve(
1446 &self,
1447 b: &TypedTensor<T>,
1448 session: &mut dyn BackendSession,
1449 ) -> tenferro_tensor::Result<TypedTensor<T>>;
1450 /// # Errors
1451 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1452 /// Returns incompatible-input validation, backend, singular, or typed-downcast errors.
1453 /// # Examples
1454 ///
1455 /// ```rust
1456 /// # use tenferro_cpu::CpuBackend;
1457 /// # use tenferro_linalg::TypedTensorLinalgExt;
1458 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1459 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1460 /// # let b = TypedTensor::<f64>::from_vec_col_major(vec![2, 1], vec![4.0, 8.0])?;
1461 /// # let mut host = CpuBackend::new();
1462 /// let x = host.with_backend_session(|session| a.solve(&b, session))??;
1463 /// assert_eq!(x.as_slice()?, &[2.0, 2.0]);
1464 /// # Ok::<(), tenferro_tensor::Error>(())
1465 /// ```
1466 fn solve(
1467 &self,
1468 b: &TypedTensor<T>,
1469 session: &mut dyn BackendSession,
1470 ) -> tenferro_tensor::Result<TypedTensor<T>>;
1471 /// # Errors
1472 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1473 /// Returns matrix validation, backend, positive-definiteness, or typed-downcast errors.
1474 /// # Examples
1475 ///
1476 /// ```rust
1477 /// # use tenferro_cpu::CpuBackend;
1478 /// # use tenferro_linalg::TypedTensorLinalgExt;
1479 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1480 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![4.0, 2.0, 2.0, 3.0])?;
1481 /// # let mut host = CpuBackend::new();
1482 /// let l = host.with_backend_session(|session| a.cholesky(session))??;
1483 /// assert_eq!(l.shape(), &[2, 2]);
1484 /// # Ok::<(), tenferro_tensor::Error>(())
1485 /// ```
1486 fn cholesky(&self, session: &mut dyn BackendSession)
1487 -> tenferro_tensor::Result<TypedTensor<T>>;
1488 /// # Errors
1489 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1490 /// Returns validation, backend, convergence, output-contract, or typed-downcast errors.
1491 /// # Examples
1492 ///
1493 /// ```rust
1494 /// # use tenferro_cpu::CpuBackend;
1495 /// # use tenferro_linalg::TypedTensorLinalgExt;
1496 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1497 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 3.0])?;
1498 /// # let mut host = CpuBackend::new();
1499 /// let (values, vectors) = host.with_backend_session(|session| a.eigh(session))??;
1500 /// assert_eq!(values.as_slice()?, &[1.0, 3.0]);
1501 /// assert_eq!(vectors.shape(), &[2, 2]);
1502 /// # Ok::<(), tenferro_tensor::Error>(())
1503 /// ```
1504 fn eigh(
1505 &self,
1506 session: &mut dyn BackendSession,
1507 ) -> tenferro_tensor::Result<(TypedTensor<<T as TensorScalar>::Real>, TypedTensor<T>)>;
1508 /// # Errors
1509 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1510 /// Returns validation errors for metadata/options, plus backend or typed-output errors.
1511 /// # Examples
1512 ///
1513 /// ```rust
1514 /// # use tenferro_cpu::CpuBackend;
1515 /// # use tenferro_linalg::{EighOptions, TypedTensorLinalgExt};
1516 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1517 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 3.0])?;
1518 /// # let mut host = CpuBackend::new();
1519 /// let (values, vectors) = host.with_backend_session(|session| { a.eigh_with_options(EighOptions::default(), session) })??;
1520 /// assert_eq!(values.shape(), &[2]);
1521 /// assert_eq!(vectors.shape(), &[2, 2]);
1522 /// # Ok::<(), tenferro_tensor::Error>(())
1523 /// ```
1524 fn eigh_with_options(
1525 &self,
1526 options: EighOptions,
1527 session: &mut dyn BackendSession,
1528 ) -> tenferro_tensor::Result<(TypedTensor<<T as TensorScalar>::Real>, TypedTensor<T>)>;
1529 /// # Errors
1530 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1531 /// Returns validation, backend, convergence, output-contract, or typed-downcast errors.
1532 /// # Examples
1533 ///
1534 /// ```rust
1535 /// # use tenferro_cpu::CpuBackend;
1536 /// # use tenferro_linalg::TypedTensorLinalgExt;
1537 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1538 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 2.0])?;
1539 /// # let mut host = CpuBackend::new();
1540 /// let (values, vectors) = host.with_backend_session(|session| a.eig(session))??;
1541 /// assert_eq!(values.shape(), &[2]);
1542 /// assert_eq!(vectors.shape(), &[2, 2]);
1543 /// # Ok::<(), tenferro_tensor::Error>(())
1544 /// ```
1545 fn eig(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedEig<T>>;
1546 /// # Errors
1547 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1548 /// Returns incompatible-input/flag validation, backend, singular, or typed-output errors.
1549 #[allow(clippy::too_many_arguments)]
1550 /// # Examples
1551 ///
1552 /// ```rust
1553 /// # use tenferro_cpu::CpuBackend;
1554 /// # use tenferro_linalg::TypedTensorLinalgExt;
1555 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1556 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 1.0, 3.0])?;
1557 /// # let b = TypedTensor::<f64>::from_vec_col_major(vec![2, 1], vec![4.0, 9.0])?;
1558 /// # let mut host = CpuBackend::new();
1559 /// let x = host.with_backend_session(|session| { a.triangular_solve(&b, true, false, false, false, session) })??;
1560 /// assert_eq!(x.shape(), &[2, 1]);
1561 /// # Ok::<(), tenferro_tensor::Error>(())
1562 /// ```
1563 fn triangular_solve(
1564 &self,
1565 b: &TypedTensor<T>,
1566 left_side: bool,
1567 lower: bool,
1568 transpose_a: bool,
1569 unit_diagonal: bool,
1570 session: &mut dyn BackendSession,
1571 ) -> tenferro_tensor::Result<TypedTensor<T>>;
1572 /// # Errors
1573 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1574 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1575 /// # Examples
1576 ///
1577 /// ```rust
1578 /// # use tenferro_cpu::CpuBackend;
1579 /// # use tenferro_linalg::TypedTensorLinalgExt;
1580 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1581 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1582 /// # let mut host = CpuBackend::new();
1583 /// let (sign, logabsdet) = host.with_backend_session(|session| a.slogdet(session))??;
1584 /// assert_eq!(sign.as_slice()?, &[1.0]);
1585 /// assert_eq!(logabsdet.shape(), &[] as &[usize]);
1586 /// # Ok::<(), tenferro_tensor::Error>(())
1587 /// ```
1588 fn slogdet(
1589 &self,
1590 session: &mut dyn BackendSession,
1591 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<<T as TensorScalar>::Real>)>;
1592 /// # Errors
1593 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1594 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1595 /// # Examples
1596 ///
1597 /// ```rust
1598 /// # use tenferro_cpu::CpuBackend;
1599 /// # use tenferro_linalg::TypedTensorLinalgExt;
1600 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1601 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1602 /// # let mut host = CpuBackend::new();
1603 /// let determinant = host.with_backend_session(|session| a.det(session))??;
1604 /// assert!((determinant.as_slice()?[0] - 8.0).abs() < 1.0e-12);
1605 /// # Ok::<(), tenferro_tensor::Error>(())
1606 /// ```
1607 fn det(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1608 /// # Errors
1609 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1610 /// Returns validation, backend, singular-solve, or typed-downcast errors.
1611 /// # Examples
1612 ///
1613 /// ```rust
1614 /// # use tenferro_cpu::CpuBackend;
1615 /// # use tenferro_linalg::TypedTensorLinalgExt;
1616 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1617 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1618 /// # let mut host = CpuBackend::new();
1619 /// let inverse = host.with_backend_session(|session| a.inv(session))??;
1620 /// assert_eq!(inverse.as_slice()?, &[0.5, 0.0, 0.0, 0.25]);
1621 /// # Ok::<(), tenferro_tensor::Error>(())
1622 /// ```
1623 fn inv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1624 /// # Errors
1625 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1626 /// Returns validation, backend, convergence, output-contract, or typed-downcast errors.
1627 /// # Examples
1628 ///
1629 /// ```rust
1630 /// # use tenferro_cpu::CpuBackend;
1631 /// # use tenferro_linalg::TypedTensorLinalgExt;
1632 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1633 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 3.0])?;
1634 /// # let mut host = CpuBackend::new();
1635 /// let values = host.with_backend_session(|session| a.eigvalsh(session))??;
1636 /// assert_eq!(values.as_slice()?, &[1.0, 3.0]);
1637 /// # Ok::<(), tenferro_tensor::Error>(())
1638 /// ```
1639 fn eigvalsh(
1640 &self,
1641 session: &mut dyn BackendSession,
1642 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>>;
1643 /// # Errors
1644 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1645 /// Returns validation, backend, convergence, output-contract, or typed-downcast errors.
1646 /// # Examples
1647 ///
1648 /// ```rust
1649 /// # use tenferro_cpu::CpuBackend;
1650 /// # use tenferro_linalg::TypedTensorLinalgExt;
1651 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1652 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 2.0])?;
1653 /// # let mut host = CpuBackend::new();
1654 /// let values = host.with_backend_session(|session| a.eigvals(session))??;
1655 /// assert_eq!(values.shape(), &[2]);
1656 /// # Ok::<(), tenferro_tensor::Error>(())
1657 /// ```
1658 fn eigvals(
1659 &self,
1660 session: &mut dyn BackendSession,
1661 ) -> tenferro_tensor::Result<TypedTensor<T::Complex>>;
1662 /// # Errors
1663 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1664 /// Returns validation, backend, numerical, output-contract, or typed-downcast errors.
1665 /// # Examples
1666 ///
1667 /// ```rust
1668 /// # use tenferro_cpu::CpuBackend;
1669 /// # use tenferro_linalg::TypedTensorLinalgExt;
1670 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1671 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1672 /// # let mut host = CpuBackend::new();
1673 /// let pseudoinverse = host.with_backend_session(|session| a.pinv(session))??;
1674 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
1675 /// # Ok::<(), tenferro_tensor::Error>(())
1676 /// ```
1677 fn pinv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1678 /// # Errors
1679 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1680 /// Returns a validation error for invalid `rtol`, plus backend or typed-output errors.
1681 /// # Examples
1682 ///
1683 /// ```rust
1684 /// # use tenferro_cpu::CpuBackend;
1685 /// # use tenferro_linalg::TypedTensorLinalgExt;
1686 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1687 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0])?;
1688 /// # let mut host = CpuBackend::new();
1689 /// let pseudoinverse = host.with_backend_session(|session| { a.pinv_with_rtol(1.0e-12, session) })??;
1690 /// assert_eq!(pseudoinverse.shape(), &[2, 2]);
1691 /// # Ok::<(), tenferro_tensor::Error>(())
1692 /// ```
1693 fn pinv_with_rtol(
1694 &self,
1695 rtol: f64,
1696 session: &mut dyn BackendSession,
1697 ) -> tenferro_tensor::Result<TypedTensor<T>>;
1698 /// # Errors
1699 /// Returns `tenferro_tensor::Error::Unsupported` when the selected backend does not support the operation.
1700 /// Returns validation errors for axes/order combinations, backend, or typed-output errors.
1701 /// # Examples
1702 ///
1703 /// ```rust
1704 /// # use tenferro_cpu::CpuBackend;
1705 /// # use tenferro_linalg::TypedTensorLinalgExt;
1706 /// # use tenferro_tensor::{BackendSessionHost, TypedTensor};
1707 /// # let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![3.0, 0.0, 0.0, 4.0])?;
1708 /// # let mut host = CpuBackend::new();
1709 /// let frobenius = host.with_backend_session(|session| a.norm(None, None, false, session))??;
1710 /// assert_eq!(frobenius.shape(), &[] as &[usize]);
1711 /// # Ok::<(), tenferro_tensor::Error>(())
1712 /// ```
1713 fn norm(
1714 &self,
1715 ord: Option<f64>,
1716 dim: Option<&[usize]>,
1717 keepdim: bool,
1718 session: &mut dyn BackendSession,
1719 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>>;
1720}
1721
1722impl TensorLinalgExt for Tensor {
1723 fn svd(
1724 &self,
1725 session: &mut dyn BackendSession,
1726 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1727 with_linalg_backend(session, "svd", |backend| three(backend.svd(self)?, "svd"))
1728 }
1729 fn svdvals(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1730 with_linalg_backend(session, "svdvals", |backend| backend.svd_values(self))
1731 }
1732
1733 fn svd_with_options(
1734 &self,
1735 options: SvdOptions,
1736 session: &mut dyn BackendSession,
1737 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1738 with_linalg_backend(session, "svd_with_options", |backend| {
1739 three(backend.svd_with_options(self, options)?, "svd_with_options")
1740 })
1741 }
1742 fn svd_full(
1743 &self,
1744 session: &mut dyn BackendSession,
1745 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1746 with_linalg_backend(session, "svd_full", |backend| {
1747 three(backend.svd_full(self)?, "svd_full")
1748 })
1749 }
1750 fn qr(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1751 with_linalg_backend(session, "qr", |backend| two(backend.qr(self)?, "qr"))
1752 }
1753 fn householder_qr(
1754 &self,
1755 session: &mut dyn BackendSession,
1756 ) -> tenferro_tensor::Result<crate::HouseholderQr<Tensor>> {
1757 with_linalg_backend(session, "householder_qr", |backend| {
1758 backend
1759 .householder_qr(self)
1760 .map(crate::HouseholderQr::from_backend)
1761 })
1762 }
1763 fn qr_with_options(
1764 &self,
1765 options: QrOptions,
1766 session: &mut dyn BackendSession,
1767 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1768 with_linalg_backend(session, "qr_with_options", |backend| {
1769 two(backend.qr_with_options(self, options)?, "qr_with_options")
1770 })
1771 }
1772 fn rank_revealing_qr(
1773 &self,
1774 options: RankRevealingQrOptions,
1775 session: &mut dyn BackendSession,
1776 ) -> tenferro_tensor::Result<RankRevealingQrResult<Tensor>> {
1777 with_linalg_backend(session, "rank_revealing_qr", |backend| {
1778 rank_revealing_qr_result(backend.rank_revealing_qr(self, options)?)
1779 })
1780 }
1781 fn lu(
1782 &self,
1783 session: &mut dyn BackendSession,
1784 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor)> {
1785 with_linalg_backend(session, "lu", |backend| four(backend.lu(self)?, "lu"))
1786 }
1787 fn full_piv_lu(
1788 &self,
1789 session: &mut dyn BackendSession,
1790 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor, Tensor)> {
1791 with_linalg_backend(session, "full_piv_lu", |backend| {
1792 five(backend.full_piv_lu(self)?, "full_piv_lu")
1793 })
1794 }
1795 fn full_piv_lu_solve(
1796 &self,
1797 b: &Tensor,
1798 session: &mut dyn BackendSession,
1799 ) -> tenferro_tensor::Result<Tensor> {
1800 with_linalg_backend(session, "full_piv_lu_solve", |backend| {
1801 backend.full_piv_lu_solve(self, b, false)
1802 })
1803 }
1804 fn solve(
1805 &self,
1806 b: &Tensor,
1807 session: &mut dyn BackendSession,
1808 ) -> tenferro_tensor::Result<Tensor> {
1809 with_linalg_backend(session, "solve", |backend| backend.solve(self, b))
1810 }
1811 fn cholesky(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1812 with_linalg_backend(session, "cholesky", |backend| backend.cholesky(self))
1813 }
1814 fn eigh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1815 with_linalg_backend(session, "eigh", |backend| two(backend.eigh(self)?, "eigh"))
1816 }
1817 fn eigh_with_options(
1818 &self,
1819 options: EighOptions,
1820 session: &mut dyn BackendSession,
1821 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1822 with_linalg_backend(session, "eigh_with_options", |backend| {
1823 two(
1824 backend.eigh_with_options(self, options)?,
1825 "eigh_with_options",
1826 )
1827 })
1828 }
1829 fn eig(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1830 with_linalg_backend(session, "eig", |backend| two(backend.eig(self)?, "eig"))
1831 }
1832 fn triangular_solve(
1833 &self,
1834 b: &Tensor,
1835 left_side: bool,
1836 lower: bool,
1837 transpose_a: bool,
1838 unit_diagonal: bool,
1839 session: &mut dyn BackendSession,
1840 ) -> tenferro_tensor::Result<Tensor> {
1841 with_linalg_backend(session, "triangular_solve", |backend| {
1842 backend.triangular_solve(self, b, left_side, lower, transpose_a, unit_diagonal)
1843 })
1844 }
1845 fn slogdet(
1846 &self,
1847 session: &mut dyn BackendSession,
1848 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1849 with_linalg_backend(session, "slogdet", |backend| {
1850 slogdet_from_lu(four(backend.lu(self)?, "slogdet")?, backend)
1851 })
1852 }
1853 fn det(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1854 with_linalg_backend(session, "det", |backend| {
1855 det_impl(self.slogdet(backend)?, backend)
1856 })
1857 }
1858 fn inv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1859 with_linalg_backend(session, "inv", |backend| inv_owned(self, backend))
1860 }
1861 fn eigvalsh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1862 with_linalg_backend(session, "eigvalsh", |backend| backend.eigh_values(self))
1863 }
1864 fn eigvals(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1865 with_linalg_backend(session, "eigvals", |backend| backend.eig_values(self))
1866 }
1867 fn pinv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1868 with_linalg_backend(session, "pinv", |backend| pinv_owned(self, None, backend))
1869 }
1870 fn pinv_with_rtol(
1871 &self,
1872 rtol: f64,
1873 session: &mut dyn BackendSession,
1874 ) -> tenferro_tensor::Result<Tensor> {
1875 with_linalg_backend(session, "pinv_with_rtol", |backend| {
1876 pinv_owned(self, Some(rtol), backend)
1877 })
1878 }
1879 fn norm(
1880 &self,
1881 ord: Option<f64>,
1882 dim: Option<&[usize]>,
1883 keepdim: bool,
1884 session: &mut dyn BackendSession,
1885 ) -> tenferro_tensor::Result<Tensor> {
1886 with_linalg_backend(session, "norm", |backend| {
1887 norm_from_read(TensorRead::from_tensor(self), ord, dim, keepdim, backend)
1888 })
1889 }
1890}
1891
1892impl TensorReadLinalgExt for TensorRead<'_> {
1893 fn svd_read(
1894 self,
1895 session: &mut dyn BackendSession,
1896 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1897 with_linalg_backend(session, "svd_read", |backend| {
1898 three(backend.svd_read(self)?, "svd_read")
1899 })
1900 }
1901 fn svdvals_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
1902 with_linalg_backend(session, "svdvals_read", |backend| {
1903 backend.svd_values_read(self)
1904 })
1905 }
1906
1907 fn svd_with_options_read(
1908 self,
1909 options: SvdOptions,
1910 session: &mut dyn BackendSession,
1911 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1912 with_linalg_backend(session, "svd_with_options_read", |backend| {
1913 three(
1914 backend.svd_with_options_read(self, options)?,
1915 "svd_with_options_read",
1916 )
1917 })
1918 }
1919 fn svd_full_read(
1920 self,
1921 session: &mut dyn BackendSession,
1922 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
1923 with_linalg_backend(session, "svd_full_read", |backend| {
1924 three(backend.svd_full_read(self)?, "svd_full_read")
1925 })
1926 }
1927 fn qr_read(
1928 self,
1929 session: &mut dyn BackendSession,
1930 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1931 with_linalg_backend(session, "qr_read", |backend| {
1932 two(backend.qr_read(self)?, "qr_read")
1933 })
1934 }
1935 fn qr_with_options_read(
1936 self,
1937 options: QrOptions,
1938 session: &mut dyn BackendSession,
1939 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
1940 with_linalg_backend(session, "qr_with_options_read", |backend| {
1941 two(
1942 backend.qr_with_options_read(self, options)?,
1943 "qr_with_options_read",
1944 )
1945 })
1946 }
1947 fn rank_revealing_qr_read(
1948 self,
1949 options: RankRevealingQrOptions,
1950 session: &mut dyn BackendSession,
1951 ) -> tenferro_tensor::Result<RankRevealingQrResult<Tensor>> {
1952 with_linalg_backend(session, "rank_revealing_qr_read", |backend| {
1953 rank_revealing_qr_result(backend.rank_revealing_qr_read(self, options)?)
1954 })
1955 }
1956 fn lu_read(
1957 self,
1958 session: &mut dyn BackendSession,
1959 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor)> {
1960 with_linalg_backend(session, "lu_read", |backend| {
1961 four(backend.lu_read(self)?, "lu_read")
1962 })
1963 }
1964 fn full_piv_lu_read(
1965 self,
1966 session: &mut dyn BackendSession,
1967 ) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor, Tensor)> {
1968 with_linalg_backend(session, "full_piv_lu_read", |backend| {
1969 five(backend.full_piv_lu_read(self)?, "full_piv_lu_read")
1970 })
1971 }
1972 fn full_piv_lu_solve_read(
1973 self,
1974 b: TensorRead<'_>,
1975 session: &mut dyn BackendSession,
1976 ) -> tenferro_tensor::Result<Tensor> {
1977 with_linalg_backend(session, "full_piv_lu_solve_read", |backend| {
1978 let b_is_vector = b.shape().len() + 1 == self.shape().len();
1979 let original_b_shape = b.shape().to_vec();
1980 let vector_as_matrix = if b_is_vector {
1981 let mut shape = vec![b.shape()[0], 1];
1982 shape.extend_from_slice(&b.shape()[1..]);
1983 Some(backend.reshape_read(b.clone(), &shape)?)
1984 } else {
1985 None
1986 };
1987 let b = vector_as_matrix.as_ref().map_or(b, TensorRead::from_tensor);
1988 let (p, l, u, q, _parity) = self.full_piv_lu_read(backend)?;
1989 let pb = linalg_matmul_read(&p, b, false, backend)?;
1990 let z = backend.triangular_solve_read(
1991 TensorRead::from_tensor(&l),
1992 TensorRead::from_tensor(&pb),
1993 true,
1994 true,
1995 false,
1996 true,
1997 )?;
1998 let w = backend.triangular_solve_read(
1999 TensorRead::from_tensor(&u),
2000 TensorRead::from_tensor(&z),
2001 true,
2002 false,
2003 false,
2004 false,
2005 )?;
2006 let mut perm: Vec<usize> = (0..q.shape().len()).collect();
2007 perm.swap(0, 1);
2008 let qt = transpose(&q, &perm, backend)?;
2009 let solution = linalg_matmul_read(&qt, TensorRead::from_tensor(&w), false, backend)?;
2010 if b_is_vector {
2011 reshape(&solution, &original_b_shape, backend)
2012 } else {
2013 Ok(solution)
2014 }
2015 })
2016 }
2017 fn solve_read(
2018 self,
2019 b: TensorRead<'_>,
2020 session: &mut dyn BackendSession,
2021 ) -> tenferro_tensor::Result<Tensor> {
2022 with_linalg_backend(session, "solve_read", |backend| backend.solve_read(self, b))
2023 }
2024 fn solve_read_into(
2025 self,
2026 b: TensorRead<'_>,
2027 out: TensorWrite<'_>,
2028 session: &mut dyn BackendSession,
2029 ) -> tenferro_tensor::Result<()> {
2030 with_linalg_backend(session, "solve_read_into", |backend| {
2031 backend.solve_read_into(self, b, out)
2032 })
2033 }
2034 fn cholesky_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2035 with_linalg_backend(session, "cholesky_read", |backend| {
2036 backend.cholesky_read(self)
2037 })
2038 }
2039 fn eigh_read(
2040 self,
2041 session: &mut dyn BackendSession,
2042 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2043 with_linalg_backend(session, "eigh_read", |backend| {
2044 two(backend.eigh_read(self)?, "eigh_read")
2045 })
2046 }
2047 fn eigh_with_options_read(
2048 self,
2049 options: EighOptions,
2050 session: &mut dyn BackendSession,
2051 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2052 with_linalg_backend(session, "eigh_with_options_read", |backend| {
2053 two(
2054 backend.eigh_with_options_read(self, options)?,
2055 "eigh_with_options_read",
2056 )
2057 })
2058 }
2059 fn eig_read(
2060 self,
2061 session: &mut dyn BackendSession,
2062 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2063 with_linalg_backend(session, "eig_read", |backend| {
2064 two(backend.eig_read(self)?, "eig_read")
2065 })
2066 }
2067 fn triangular_solve_read(
2068 self,
2069 b: TensorRead<'_>,
2070 left_side: bool,
2071 lower: bool,
2072 transpose_a: bool,
2073 unit_diagonal: bool,
2074 session: &mut dyn BackendSession,
2075 ) -> tenferro_tensor::Result<Tensor> {
2076 with_linalg_backend(session, "triangular_solve_read", |backend| {
2077 backend.triangular_solve_read(self, b, left_side, lower, transpose_a, unit_diagonal)
2078 })
2079 }
2080 fn slogdet_read(
2081 self,
2082 session: &mut dyn BackendSession,
2083 ) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2084 with_linalg_backend(session, "slogdet_read", |backend| {
2085 slogdet_from_lu(four(backend.lu_read(self)?, "slogdet_read")?, backend)
2086 })
2087 }
2088 fn det_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2089 with_linalg_backend(session, "det_read", |backend| {
2090 det_impl(self.slogdet_read(backend)?, backend)
2091 })
2092 }
2093 fn inv_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2094 with_linalg_backend(session, "inv_read", |backend| inv_read(self, backend))
2095 }
2096 fn eigvalsh_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2097 with_linalg_backend(session, "eigvalsh_read", |backend| {
2098 backend.eigh_values_read(self)
2099 })
2100 }
2101 fn eigvals_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2102 with_linalg_backend(session, "eigvals_read", |backend| {
2103 backend.eig_values_read(self)
2104 })
2105 }
2106 fn pinv_read(self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor> {
2107 with_linalg_backend(session, "pinv_read", |backend| {
2108 pinv_read(self, None, backend)
2109 })
2110 }
2111 fn pinv_with_rtol_read(
2112 self,
2113 rtol: f64,
2114 session: &mut dyn BackendSession,
2115 ) -> tenferro_tensor::Result<Tensor> {
2116 with_linalg_backend(session, "pinv_with_rtol_read", |backend| {
2117 pinv_read(self, Some(rtol), backend)
2118 })
2119 }
2120 fn norm_read(
2121 self,
2122 ord: Option<f64>,
2123 dim: Option<&[usize]>,
2124 keepdim: bool,
2125 session: &mut dyn BackendSession,
2126 ) -> tenferro_tensor::Result<Tensor> {
2127 with_linalg_backend(session, "norm_read", |backend| {
2128 norm_from_read(self, ord, dim, keepdim, backend)
2129 })
2130 }
2131}
2132
2133impl<T: LinalgScalar> TypedTensorLinalgExt<T> for TypedTensor<T> {
2134 fn svd(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedSvd<T>> {
2135 with_linalg_backend(session, "svd", |backend| {
2136 typed_svd(T::tensor_read(self).svd_read(backend)?)
2137 })
2138 }
2139
2140 fn svdvals(
2141 &self,
2142 session: &mut dyn BackendSession,
2143 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>> {
2144 with_linalg_backend(session, "svdvals", |backend| {
2145 typed_output::<<T as TensorScalar>::Real>(T::tensor_read(self).svdvals_read(backend)?)
2146 })
2147 }
2148
2149 fn svd_with_options(
2150 &self,
2151 options: SvdOptions,
2152 session: &mut dyn BackendSession,
2153 ) -> tenferro_tensor::Result<TypedSvd<T>> {
2154 with_linalg_backend(session, "svd_with_options", |backend| {
2155 typed_svd(T::tensor_read(self).svd_with_options_read(options, backend)?)
2156 })
2157 }
2158
2159 fn svd_full(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedSvd<T>> {
2160 with_linalg_backend(session, "svd_full", |backend| {
2161 typed_svd(T::tensor_read(self).svd_full_read(backend)?)
2162 })
2163 }
2164
2165 fn qr(
2166 &self,
2167 session: &mut dyn BackendSession,
2168 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<T>)> {
2169 with_linalg_backend(session, "qr", |backend| {
2170 typed_pair_same(T::tensor_read(self).qr_read(backend)?)
2171 })
2172 }
2173
2174 fn qr_with_options(
2175 &self,
2176 options: QrOptions,
2177 session: &mut dyn BackendSession,
2178 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<T>)> {
2179 with_linalg_backend(session, "qr_with_options", |backend| {
2180 typed_pair_same(T::tensor_read(self).qr_with_options_read(options, backend)?)
2181 })
2182 }
2183
2184 fn rank_revealing_qr(
2185 &self,
2186 options: RankRevealingQrOptions,
2187 session: &mut dyn BackendSession,
2188 ) -> tenferro_tensor::Result<TypedRankRevealingQrResult<T>> {
2189 with_linalg_backend(session, "rank_revealing_qr", |backend| {
2190 let result = T::tensor_read(self).rank_revealing_qr_read(options, backend)?;
2191 Ok(RankRevealingQrResult {
2192 q: typed_output::<T>(result.q)?,
2193 r: typed_output::<T>(result.r)?,
2194 column_permutation: typed_output::<i64>(result.column_permutation)?,
2195 rank: typed_output::<i64>(result.rank)?,
2196 })
2197 })
2198 }
2199
2200 fn lu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedLu<T>> {
2201 with_linalg_backend(session, "lu", |backend| {
2202 let (p, l, u, parity) = T::tensor_read(self).lu_read(backend)?;
2203 Ok((
2204 typed_output::<T>(p)?,
2205 typed_output::<T>(l)?,
2206 typed_output::<T>(u)?,
2207 typed_output::<<T as TensorScalar>::Real>(parity)?,
2208 ))
2209 })
2210 }
2211
2212 fn full_piv_lu_solve(
2213 &self,
2214 b: &TypedTensor<T>,
2215 session: &mut dyn BackendSession,
2216 ) -> tenferro_tensor::Result<TypedTensor<T>> {
2217 with_linalg_backend(session, "full_piv_lu_solve", |backend| {
2218 typed_output::<T>(
2219 T::tensor_read(self).full_piv_lu_solve_read(T::tensor_read(b), backend)?,
2220 )
2221 })
2222 }
2223
2224 fn full_piv_lu(
2225 &self,
2226 session: &mut dyn BackendSession,
2227 ) -> tenferro_tensor::Result<TypedFullPivLu<T>> {
2228 with_linalg_backend(session, "full_piv_lu", |backend| {
2229 let (p, l, u, q, parity) = T::tensor_read(self).full_piv_lu_read(backend)?;
2230 Ok((
2231 typed_output::<T>(p)?,
2232 typed_output::<T>(l)?,
2233 typed_output::<T>(u)?,
2234 typed_output::<T>(q)?,
2235 typed_output::<<T as TensorScalar>::Real>(parity)?,
2236 ))
2237 })
2238 }
2239
2240 fn solve(
2241 &self,
2242 b: &TypedTensor<T>,
2243 session: &mut dyn BackendSession,
2244 ) -> tenferro_tensor::Result<TypedTensor<T>> {
2245 with_linalg_backend(session, "solve", |backend| {
2246 typed_output::<T>(T::tensor_read(self).solve_read(T::tensor_read(b), backend)?)
2247 })
2248 }
2249
2250 fn cholesky(
2251 &self,
2252 session: &mut dyn BackendSession,
2253 ) -> tenferro_tensor::Result<TypedTensor<T>> {
2254 with_linalg_backend(session, "cholesky", |backend| {
2255 typed_output::<T>(T::tensor_read(self).cholesky_read(backend)?)
2256 })
2257 }
2258
2259 fn eigh(
2260 &self,
2261 session: &mut dyn BackendSession,
2262 ) -> tenferro_tensor::Result<(TypedTensor<<T as TensorScalar>::Real>, TypedTensor<T>)> {
2263 with_linalg_backend(session, "eigh", |backend| {
2264 typed_eigh(T::tensor_read(self).eigh_read(backend)?)
2265 })
2266 }
2267
2268 fn eigh_with_options(
2269 &self,
2270 options: EighOptions,
2271 session: &mut dyn BackendSession,
2272 ) -> tenferro_tensor::Result<(TypedTensor<<T as TensorScalar>::Real>, TypedTensor<T>)> {
2273 with_linalg_backend(session, "eigh_with_options", |backend| {
2274 typed_eigh(T::tensor_read(self).eigh_with_options_read(options, backend)?)
2275 })
2276 }
2277
2278 fn eig(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedEig<T>> {
2279 with_linalg_backend(session, "eig", |backend| {
2280 let (values, vectors) = T::tensor_read(self).eig_read(backend)?;
2281 Ok((
2282 typed_output::<T::Complex>(values)?,
2283 typed_output::<T::Complex>(vectors)?,
2284 ))
2285 })
2286 }
2287
2288 fn triangular_solve(
2289 &self,
2290 b: &TypedTensor<T>,
2291 left_side: bool,
2292 lower: bool,
2293 transpose_a: bool,
2294 unit_diagonal: bool,
2295 session: &mut dyn BackendSession,
2296 ) -> tenferro_tensor::Result<TypedTensor<T>> {
2297 with_linalg_backend(session, "triangular_solve", |backend| {
2298 typed_output::<T>(T::tensor_read(self).triangular_solve_read(
2299 T::tensor_read(b),
2300 left_side,
2301 lower,
2302 transpose_a,
2303 unit_diagonal,
2304 backend,
2305 )?)
2306 })
2307 }
2308
2309 fn slogdet(
2310 &self,
2311 session: &mut dyn BackendSession,
2312 ) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<<T as TensorScalar>::Real>)> {
2313 with_linalg_backend(session, "slogdet", |backend| {
2314 let (sign, logabsdet) = T::tensor_read(self).slogdet_read(backend)?;
2315 Ok((
2316 typed_output::<T>(sign)?,
2317 typed_output::<<T as TensorScalar>::Real>(logabsdet)?,
2318 ))
2319 })
2320 }
2321
2322 fn det(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>> {
2323 with_linalg_backend(session, "det", |backend| {
2324 typed_output::<T>(T::tensor_read(self).det_read(backend)?)
2325 })
2326 }
2327
2328 fn inv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>> {
2329 with_linalg_backend(session, "inv", |backend| {
2330 typed_output::<T>(T::tensor_read(self).inv_read(backend)?)
2331 })
2332 }
2333
2334 fn eigvalsh(
2335 &self,
2336 session: &mut dyn BackendSession,
2337 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>> {
2338 with_linalg_backend(session, "eigvalsh", |backend| {
2339 typed_output::<<T as TensorScalar>::Real>(T::tensor_read(self).eigvalsh_read(backend)?)
2340 })
2341 }
2342
2343 fn eigvals(
2344 &self,
2345 session: &mut dyn BackendSession,
2346 ) -> tenferro_tensor::Result<TypedTensor<T::Complex>> {
2347 with_linalg_backend(session, "eigvals", |backend| {
2348 typed_output::<T::Complex>(T::tensor_read(self).eigvals_read(backend)?)
2349 })
2350 }
2351
2352 fn pinv(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>> {
2353 with_linalg_backend(session, "pinv", |backend| {
2354 typed_output::<T>(T::tensor_read(self).pinv_read(backend)?)
2355 })
2356 }
2357
2358 fn pinv_with_rtol(
2359 &self,
2360 rtol: f64,
2361 session: &mut dyn BackendSession,
2362 ) -> tenferro_tensor::Result<TypedTensor<T>> {
2363 with_linalg_backend(session, "pinv_with_rtol", |backend| {
2364 typed_output::<T>(T::tensor_read(self).pinv_with_rtol_read(rtol, backend)?)
2365 })
2366 }
2367
2368 fn norm(
2369 &self,
2370 ord: Option<f64>,
2371 dim: Option<&[usize]>,
2372 keepdim: bool,
2373 session: &mut dyn BackendSession,
2374 ) -> tenferro_tensor::Result<TypedTensor<<T as TensorScalar>::Real>> {
2375 with_linalg_backend(session, "norm", |backend| {
2376 typed_output::<<T as TensorScalar>::Real>(
2377 T::tensor_read(self).norm_read(ord, dim, keepdim, backend)?,
2378 )
2379 })
2380 }
2381}
2382
2383fn rank_revealing_qr_result(
2384 outputs: Vec<Tensor>,
2385) -> tenferro_tensor::Result<RankRevealingQrResult<Tensor>> {
2386 let (q, r, column_permutation, rank) = four(outputs, "rank_revealing_qr")?;
2387 Ok(RankRevealingQrResult {
2388 q,
2389 r,
2390 column_permutation,
2391 rank,
2392 })
2393}
2394
2395fn typed_svd<T: LinalgScalar>(
2396 (u, s, vt): (Tensor, Tensor, Tensor),
2397) -> tenferro_tensor::Result<TypedSvd<T>> {
2398 Ok((
2399 typed_output::<T>(u)?,
2400 typed_output::<<T as TensorScalar>::Real>(s)?,
2401 typed_output::<T>(vt)?,
2402 ))
2403}
2404
2405fn typed_pair_same<T: LinalgScalar>(
2406 (a, b): (Tensor, Tensor),
2407) -> tenferro_tensor::Result<(TypedTensor<T>, TypedTensor<T>)> {
2408 Ok((typed_output::<T>(a)?, typed_output::<T>(b)?))
2409}
2410
2411fn typed_eigh<T: LinalgScalar>(
2412 (values, vectors): (Tensor, Tensor),
2413) -> tenferro_tensor::Result<(TypedTensor<<T as TensorScalar>::Real>, TypedTensor<T>)> {
2414 Ok((
2415 typed_output::<<T as TensorScalar>::Real>(values)?,
2416 typed_output::<T>(vectors)?,
2417 ))
2418}
2419
2420/// Run a concrete linalg body against the built-in linalg execution sessions
2421/// (CPU/CUDA) carried by `session`, returning a typed capability error when
2422/// the session does not expose a linalg execution capability.
2423///
2424/// This is the built-in dispatch shared by the concrete linalg surface;
2425/// callers never downcast themselves (issue #1680 Phase 3). Third-party
2426/// [`LinalgBackend`] implementations remain supported through the SPI trait,
2427/// but the concrete op path is built-in-session only. The composite bodies
2428/// run on the borrowed `&mut dyn LinalgBackend` exactly as before.
2429pub(crate) fn with_linalg_backend<X>(
2430 session: &mut dyn BackendSession,
2431 op: &'static str,
2432 f: impl FnOnce(&mut dyn LinalgBackend) -> tenferro_tensor::Result<X>,
2433) -> tenferro_tensor::Result<X> {
2434 // The capability branches are mutually exclusive, so `f` runs exactly
2435 // once. Probe the marker first, then re-extract the same exec session and
2436 // run the composite body on it (FnOnce cannot be captured by several
2437 // branch closures).
2438 if with_cpu_exec_session(session, |_| ()).is_some() {
2439 return with_cpu_exec_session(session, |exec| f(exec as &mut dyn LinalgBackend))
2440 .expect("marker probe matched a CPU execution session");
2441 }
2442 #[cfg(feature = "cuda")]
2443 if with_cuda_exec_session(session, |_| ()).is_some() {
2444 return with_cuda_exec_session(session, |exec| f(exec as &mut dyn LinalgBackend))
2445 .expect("marker probe matched a CUDA execution session");
2446 }
2447 Err(tenferro_tensor::Error::unsupported(
2448 op,
2449 "selected backend session does not expose a linalg execution capability",
2450 ))
2451}
2452
2453fn typed_output<T: TensorScalar>(tensor: Tensor) -> tenferro_tensor::Result<TypedTensor<T>> {
2454 if tensor.dtype() != T::dtype() {
2455 return Err(tenferro_tensor::Error::Internal(format!(
2456 "typed linalg backend contract expected {:?}, got {:?}",
2457 T::dtype(),
2458 tensor.dtype()
2459 )));
2460 }
2461 match T::into_typed(tensor) {
2462 Ok(typed) => Ok(typed),
2463 // INVARIANT: the dtype guard above already accepted this tensor, so the
2464 // refusal arm is unreachable for a matching preset scalar.
2465 Err(failure) => Err(tenferro_tensor::Error::Internal(format!(
2466 "typed linalg backend contract lost its dtype guard: {}",
2467 failure.error()
2468 ))),
2469 }
2470}
2471
2472fn arity(name: &'static str, expected: usize, actual: usize) -> tenferro_tensor::Error {
2473 tenferro_tensor::Error::Internal(format!(
2474 "{name} backend contract expected {expected} outputs, got {actual}"
2475 ))
2476}
2477fn two(mut out: Vec<Tensor>, name: &'static str) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2478 if out.len() != 2 {
2479 return Err(arity(name, 2, out.len()));
2480 }
2481 let b = out.pop().unwrap();
2482 let a = out.pop().unwrap();
2483 Ok((a, b))
2484}
2485fn three(
2486 mut out: Vec<Tensor>,
2487 name: &'static str,
2488) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor)> {
2489 if out.len() != 3 {
2490 return Err(arity(name, 3, out.len()));
2491 }
2492 let c = out.pop().unwrap();
2493 let b = out.pop().unwrap();
2494 let a = out.pop().unwrap();
2495 Ok((a, b, c))
2496}
2497fn four(
2498 mut out: Vec<Tensor>,
2499 name: &'static str,
2500) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor)> {
2501 if out.len() != 4 {
2502 return Err(arity(name, 4, out.len()));
2503 }
2504 let d = out.pop().unwrap();
2505 let c = out.pop().unwrap();
2506 let b = out.pop().unwrap();
2507 let a = out.pop().unwrap();
2508 Ok((a, b, c, d))
2509}
2510fn five(
2511 mut out: Vec<Tensor>,
2512 name: &'static str,
2513) -> tenferro_tensor::Result<(Tensor, Tensor, Tensor, Tensor, Tensor)> {
2514 if out.len() != 5 {
2515 return Err(arity(name, 5, out.len()));
2516 }
2517 let e = out.pop().unwrap();
2518 let d = out.pop().unwrap();
2519 let c = out.pop().unwrap();
2520 let b = out.pop().unwrap();
2521 let a = out.pop().unwrap();
2522 Ok((a, b, c, d, e))
2523}
2524
2525// Composite implementations are kept below the surface adapters so every
2526// backend-visible operation remains explicit and testable.
2527fn slogdet_from_lu<B: LinalgBackend + ?Sized>(
2528 (_p, _l, u, parity): (Tensor, Tensor, Tensor, Tensor),
2529 backend: &mut B,
2530) -> tenferro_tensor::Result<(Tensor, Tensor)> {
2531 let diag = backend.extract_diagonal(&u, 0, 1)?;
2532 let sign = backend.sign_read(TensorRead::from_tensor(&diag))?;
2533 let sign_u = backend.reduce_prod_read(TensorRead::from_tensor(&sign), &[0])?;
2534 let sign = backend.mul_read(
2535 TensorRead::from_tensor(&parity),
2536 TensorRead::from_tensor(&sign_u),
2537 )?;
2538 let abs = backend.abs_read(TensorRead::from_tensor(&diag))?;
2539 let log = backend.log_read(TensorRead::from_tensor(&abs))?;
2540 let logabsdet = backend.reduce_sum_read(TensorRead::from_tensor(&log), &[0])?;
2541 Ok((sign, logabsdet))
2542}
2543fn det_impl<B: LinalgBackend + ?Sized>(
2544 (sign, logabsdet): (Tensor, Tensor),
2545 backend: &mut B,
2546) -> tenferro_tensor::Result<Tensor> {
2547 let magnitude = backend.exp_read(TensorRead::from_tensor(&logabsdet))?;
2548 backend.mul_read(
2549 TensorRead::from_tensor(&sign),
2550 TensorRead::from_tensor(&magnitude),
2551 )
2552}
2553fn inv_owned<B: LinalgBackend + ?Sized>(
2554 a: &Tensor,
2555 backend: &mut B,
2556) -> tenferro_tensor::Result<Tensor> {
2557 let eye = eye_like(a.dtype(), a.shape(), backend)?;
2558 backend.solve(a, &eye)
2559}
2560fn inv_read<B: LinalgBackend + ?Sized>(
2561 a: TensorRead<'_>,
2562 backend: &mut B,
2563) -> tenferro_tensor::Result<Tensor> {
2564 let eye = eye_like(a.dtype(), a.shape(), backend)?;
2565 backend.solve_read(a, TensorRead::from_tensor(&eye))
2566}
2567fn pinv_owned<B: LinalgBackend + ?Sized>(
2568 a: &Tensor,
2569 rtol: Option<f64>,
2570 backend: &mut B,
2571) -> tenferro_tensor::Result<Tensor> {
2572 ensure_float_or_complex("pinv", a.dtype())?;
2573 let outputs = three(backend.svd(a)?, "pinv")?;
2574 pinv_from_svd(
2575 outputs,
2576 rtol.unwrap_or_else(|| default_pinv_rtol(a.dtype(), a.shape())),
2577 backend,
2578 )
2579}
2580fn pinv_read<B: LinalgBackend + ?Sized>(
2581 a: TensorRead<'_>,
2582 rtol: Option<f64>,
2583 backend: &mut B,
2584) -> tenferro_tensor::Result<Tensor> {
2585 ensure_float_or_complex("pinv", a.dtype())?;
2586 let default_rtol = default_pinv_rtol(a.dtype(), a.shape());
2587 let outputs = three(backend.svd_read(a)?, "pinv_read")?;
2588 pinv_from_svd(outputs, rtol.unwrap_or(default_rtol), backend)
2589}
2590fn norm_from_read<B: LinalgBackend + ?Sized>(
2591 input: TensorRead<'_>,
2592 ord: Option<f64>,
2593 dim: Option<&[usize]>,
2594 keepdim: bool,
2595 backend: &mut B,
2596) -> tenferro_tensor::Result<Tensor> {
2597 ensure_float_or_complex("norm", input.dtype())?;
2598 let original_shape = input.shape().to_vec();
2599 let axes = dim.map_or_else(
2600 || (0..original_shape.len()).collect::<Vec<_>>(),
2601 <[usize]>::to_vec,
2602 );
2603 validate_axes("norm", original_shape.len(), &axes)?;
2604 if axes.is_empty() {
2605 return backend.to_contiguous_read(input);
2606 }
2607 let reduced = if can_square_without_abs(input.dtype(), axes.len(), ord) {
2608 frobenius_norm_read(input.clone(), &axes, backend)?
2609 } else if axes.len() == 2 {
2610 matrix_norm(input, &axes, ord, backend)?
2611 } else {
2612 let abs = backend.abs_read(input)?;
2613 norm_over_axes(&abs, &axes, ord, backend)?
2614 };
2615 if !keepdim {
2616 return Ok(reduced);
2617 }
2618 let mut shape = original_shape;
2619 for &axis in &axes {
2620 shape[axis] = 1;
2621 }
2622 reshape(&reduced, &shape, backend)
2623}
2624
2625fn scalar_real(dtype: DType, value: f64) -> tenferro_tensor::Result<Tensor> {
2626 match dtype {
2627 DType::F32 => Tensor::from_vec_col_major(vec![], vec![value as f32]),
2628 DType::F64 => Tensor::from_vec_col_major(vec![], vec![value]),
2629 DType::C32 => Tensor::from_vec_col_major(vec![], vec![Complex32::new(value as f32, 0.0)]),
2630 DType::C64 => Tensor::from_vec_col_major(vec![], vec![Complex64::new(value, 0.0)]),
2631 _ => Err(crate::error::unsupported_dtype("linalg_scalar", dtype)),
2632 }
2633}
2634
2635fn broadcast<B: LinalgBackend + ?Sized>(
2636 input: &Tensor,
2637 shape: &[usize],
2638 dims: &[usize],
2639 backend: &mut B,
2640) -> tenferro_tensor::Result<Tensor> {
2641 if input.shape() == shape {
2642 return input.duplicate();
2643 }
2644 backend.broadcast_in_dim_read(TensorRead::from_tensor(input), shape, dims)
2645}
2646
2647fn eye_like<B: LinalgBackend + ?Sized>(
2648 dtype: DType,
2649 shape: &[usize],
2650 backend: &mut B,
2651) -> tenferro_tensor::Result<Tensor> {
2652 if shape.len() < 2 {
2653 return Err(tenferro_tensor::Error::rank_mismatch("inv", 2, shape.len()));
2654 }
2655 let mut diagonal_shape = vec![shape[0]];
2656 diagonal_shape.extend_from_slice(&shape[2..]);
2657 let scalar = scalar_real(dtype, 1.0)?;
2658 let diagonal = broadcast(&scalar, &diagonal_shape, &[], backend)?;
2659 backend.embed_diagonal(&diagonal, 0, 1)
2660}
2661
2662fn ensure_float_or_complex(op: &'static str, dtype: DType) -> tenferro_tensor::Result<()> {
2663 match dtype {
2664 DType::F32 | DType::F64 | DType::C32 | DType::C64 => Ok(()),
2665 _ => Err(crate::error::unsupported_dtype(op, dtype)),
2666 }
2667}
2668
2669fn can_square_without_abs(dtype: DType, axes_len: usize, ord: Option<f64>) -> bool {
2670 matches!(dtype, DType::F32 | DType::F64)
2671 && (ord.is_none() || (ord == Some(2.0) && axes_len != 2))
2672}
2673
2674fn validate_axes(op: &'static str, rank: usize, axes: &[usize]) -> tenferro_tensor::Result<()> {
2675 tenferro_tensor::validate::validate_unique_axes(op, "dim", rank, axes)
2676}
2677
2678fn default_pinv_rtol(dtype: DType, shape: &[usize]) -> f64 {
2679 let max_dim = shape
2680 .first()
2681 .copied()
2682 .unwrap_or(0)
2683 .max(shape.get(1).copied().unwrap_or(0));
2684 let eps = match dtype {
2685 DType::F32 | DType::C32 => f32::EPSILON as f64,
2686 DType::F64 | DType::C64 => f64::EPSILON,
2687 _ => 0.0,
2688 };
2689 eps * max_dim as f64
2690}
2691
2692fn pinv_from_svd<B: LinalgBackend + ?Sized>(
2693 (u, s, vt): (Tensor, Tensor, Tensor),
2694 rtol: f64,
2695 backend: &mut B,
2696) -> tenferro_tensor::Result<Tensor> {
2697 let abs_s = backend.abs_read(TensorRead::from_tensor(&s))?;
2698 let s_max = reduce_max(&abs_s, &[0], backend)?;
2699 let threshold_scalar = scalar_real(abs_s.dtype(), rtol.max(0.0))?;
2700 let threshold = binary_mul(&s_max, &threshold_scalar, backend)?;
2701 let threshold = broadcast_batch_scalar_to_leading_axis(&threshold, s.shape(), backend)?;
2702 let mask = {
2703 let mask = backend.compare_read(
2704 TensorRead::from_tensor(&abs_s),
2705 TensorRead::from_tensor(&threshold),
2706 &CompareDir::Gt,
2707 )?;
2708 backend.convert(&mask, s.dtype())?
2709 };
2710 let ones = ones_like(&s, backend)?;
2711 let neg_mask = backend.neg_read(TensorRead::from_tensor(&mask))?;
2712 let offset = binary_add(&ones, &neg_mask, backend)?;
2713 let denom = binary_add(&s, &offset, backend)?;
2714 let s_inv = backend.div_read(
2715 TensorRead::from_tensor(&mask),
2716 TensorRead::from_tensor(&denom),
2717 )?;
2718 let v = conjugate_transpose(&vt, backend)?;
2719 let uh = conjugate_transpose(&u, backend)?;
2720 let vs = scale_matrix_columns(&v, &s_inv, backend)?;
2721 matmul_preserve_trailing_batch(&vs, &uh, backend)
2722}
2723
2724fn ones_like<B: LinalgBackend + ?Sized>(
2725 input: &Tensor,
2726 backend: &mut B,
2727) -> tenferro_tensor::Result<Tensor> {
2728 let one = scalar_real(input.dtype(), 1.0)?;
2729 broadcast(&one, input.shape(), &[], backend)
2730}
2731
2732fn binary_add<B: LinalgBackend + ?Sized>(
2733 lhs: &Tensor,
2734 rhs: &Tensor,
2735 backend: &mut B,
2736) -> tenferro_tensor::Result<Tensor> {
2737 backend.add_read(TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs))
2738}
2739
2740fn binary_mul<B: LinalgBackend + ?Sized>(
2741 lhs: &Tensor,
2742 rhs: &Tensor,
2743 backend: &mut B,
2744) -> tenferro_tensor::Result<Tensor> {
2745 backend.mul_read(TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs))
2746}
2747
2748fn reduce_max<B: LinalgBackend + ?Sized>(
2749 input: &Tensor,
2750 axes: &[usize],
2751 backend: &mut B,
2752) -> tenferro_tensor::Result<Tensor> {
2753 backend.reduce_max_read(TensorRead::from_tensor(input), axes)
2754}
2755
2756fn reduce_sum<B: LinalgBackend + ?Sized>(
2757 input: &Tensor,
2758 axes: &[usize],
2759 backend: &mut B,
2760) -> tenferro_tensor::Result<Tensor> {
2761 backend.reduce_sum_read(TensorRead::from_tensor(input), axes)
2762}
2763
2764fn reduce_min<B: LinalgBackend + ?Sized>(
2765 input: &Tensor,
2766 axes: &[usize],
2767 backend: &mut B,
2768) -> tenferro_tensor::Result<Tensor> {
2769 backend.reduce_min_read(TensorRead::from_tensor(input), axes)
2770}
2771
2772fn reshape<B: LinalgBackend + ?Sized>(
2773 input: &Tensor,
2774 shape: &[usize],
2775 backend: &mut B,
2776) -> tenferro_tensor::Result<Tensor> {
2777 backend.reshape_read(TensorRead::from_tensor(input), shape)
2778}
2779
2780fn transpose<B: LinalgBackend + ?Sized>(
2781 input: &Tensor,
2782 perm: &[usize],
2783 backend: &mut B,
2784) -> tenferro_tensor::Result<Tensor> {
2785 backend.transpose_read(TensorRead::from_tensor(input), perm)
2786}
2787
2788fn broadcast_batch_scalar_to_leading_axis<B: LinalgBackend + ?Sized>(
2789 input: &Tensor,
2790 shape: &[usize],
2791 backend: &mut B,
2792) -> tenferro_tensor::Result<Tensor> {
2793 let dims: Vec<usize> = (1..shape.len()).collect();
2794 broadcast(input, shape, &dims, backend)
2795}
2796
2797fn conjugate_transpose<B: LinalgBackend + ?Sized>(
2798 input: &Tensor,
2799 backend: &mut B,
2800) -> tenferro_tensor::Result<Tensor> {
2801 let conj = backend.conj_read(TensorRead::from_tensor(input))?;
2802 let mut perm: Vec<usize> = (0..input.shape().len()).collect();
2803 perm.swap(0, 1);
2804 transpose(&conj, &perm, backend)
2805}
2806
2807fn scale_matrix_columns<B: LinalgBackend + ?Sized>(
2808 matrix: &Tensor,
2809 scale: &Tensor,
2810 backend: &mut B,
2811) -> tenferro_tensor::Result<Tensor> {
2812 let mut scale_shape = vec![1, scale.shape()[0]];
2813 scale_shape.extend_from_slice(&matrix.shape()[2..]);
2814 let reshaped = reshape(scale, &scale_shape, backend)?;
2815 let dims: Vec<usize> = (0..matrix.shape().len()).collect();
2816 let expanded = broadcast(&reshaped, matrix.shape(), &dims, backend)?;
2817 binary_mul(matrix, &expanded, backend)
2818}
2819
2820fn matmul_preserve_trailing_batch<B: BackendSession + ?Sized>(
2821 lhs: &Tensor,
2822 rhs: &Tensor,
2823 session: &mut B,
2824) -> tenferro_tensor::Result<Tensor> {
2825 let batch: Vec<usize> = (2..lhs.shape().len()).collect();
2826 let config = DotGeneralConfig {
2827 lhs_contracting_dims: [1].as_slice().into(),
2828 rhs_contracting_dims: [0].as_slice().into(),
2829 lhs_batch_dims: batch.clone().into(),
2830 rhs_batch_dims: batch.into(),
2831 };
2832 session.dot_general_read(
2833 TensorRead::from_tensor(lhs),
2834 TensorRead::from_tensor(rhs),
2835 &config,
2836 )
2837}
2838
2839fn linalg_matmul_read<B: BackendSession + ?Sized>(
2840 lhs: &Tensor,
2841 rhs: TensorRead<'_>,
2842 rhs_is_vector: bool,
2843 session: &mut B,
2844) -> tenferro_tensor::Result<Tensor> {
2845 let lhs_batch_dims: Vec<usize> = (2..lhs.shape().len()).collect();
2846 let rhs_batch_start = if rhs_is_vector { 1 } else { 2 };
2847 let rhs_batch_dims: Vec<usize> = (rhs_batch_start..rhs.shape().len()).collect();
2848 let config = DotGeneralConfig {
2849 lhs_contracting_dims: [1].as_slice().into(),
2850 rhs_contracting_dims: [0].as_slice().into(),
2851 lhs_batch_dims: lhs_batch_dims.into(),
2852 rhs_batch_dims: rhs_batch_dims.into(),
2853 };
2854 session.dot_general_read(TensorRead::from_tensor(lhs), rhs, &config)
2855}
2856
2857fn frobenius_norm<B: LinalgBackend + ?Sized>(
2858 abs: &Tensor,
2859 axes: &[usize],
2860 backend: &mut B,
2861) -> tenferro_tensor::Result<Tensor> {
2862 let sum = backend.reduce_sum_squares_read(TensorRead::from_tensor(abs), axes)?;
2863 backend.sqrt_read(TensorRead::from_tensor(&sum))
2864}
2865
2866fn frobenius_norm_read<B: LinalgBackend + ?Sized>(
2867 input: TensorRead<'_>,
2868 axes: &[usize],
2869 backend: &mut B,
2870) -> tenferro_tensor::Result<Tensor> {
2871 let sum = backend.reduce_sum_squares_read(input, axes)?;
2872 backend.sqrt_read(TensorRead::from_tensor(&sum))
2873}
2874
2875fn p_norm<B: LinalgBackend + ?Sized>(
2876 abs: &Tensor,
2877 axes: &[usize],
2878 p: f64,
2879 backend: &mut B,
2880) -> tenferro_tensor::Result<Tensor> {
2881 if !p.is_finite() || p == 0.0 {
2882 return Err(tenferro_tensor::Error::invalid_argument(
2883 "norm",
2884 "p",
2885 format!("p-norm order must be finite and nonzero, got {p}"),
2886 ));
2887 }
2888 if p == 2.0 {
2889 return frobenius_norm(abs, axes, backend);
2890 }
2891 let power = scalar_real(abs.dtype(), p)?;
2892 let powered = backend.pow_read(
2893 TensorRead::from_tensor(abs),
2894 TensorRead::from_tensor(&power),
2895 )?;
2896 let sum = reduce_sum(&powered, axes, backend)?;
2897 let inverse = scalar_real(abs.dtype(), 1.0 / p)?;
2898 backend.pow_read(
2899 TensorRead::from_tensor(&sum),
2900 TensorRead::from_tensor(&inverse),
2901 )
2902}
2903
2904fn count_nonzero<B: LinalgBackend + ?Sized>(
2905 abs: &Tensor,
2906 axes: &[usize],
2907 backend: &mut B,
2908) -> tenferro_tensor::Result<Tensor> {
2909 let zero = scalar_real(abs.dtype(), 0.0)?;
2910 let zero = broadcast(&zero, abs.shape(), &[], backend)?;
2911 let mask = backend.compare_read(
2912 TensorRead::from_tensor(abs),
2913 TensorRead::from_tensor(&zero),
2914 &CompareDir::Gt,
2915 )?;
2916 let mask = backend.convert(&mask, abs.dtype())?;
2917 reduce_sum(&mask, axes, backend)
2918}
2919
2920fn norm_over_axes<B: LinalgBackend + ?Sized>(
2921 abs: &Tensor,
2922 axes: &[usize],
2923 ord: Option<f64>,
2924 backend: &mut B,
2925) -> tenferro_tensor::Result<Tensor> {
2926 match ord {
2927 None => frobenius_norm(abs, axes, backend),
2928 Some(p) if p == f64::INFINITY => reduce_max(abs, axes, backend),
2929 Some(p) if p == f64::NEG_INFINITY => reduce_min(abs, axes, backend),
2930 Some(0.0) => count_nonzero(abs, axes, backend),
2931 Some(p) => p_norm(abs, axes, p, backend),
2932 }
2933}
2934
2935fn matrix_norm<B: LinalgBackend + ?Sized>(
2936 input: TensorRead<'_>,
2937 axes: &[usize],
2938 ord: Option<f64>,
2939 backend: &mut B,
2940) -> tenferro_tensor::Result<Tensor> {
2941 let matrix = move_axes_to_front(input, axes, backend)?;
2942 if matches!(ord, Some(2.0) | Some(-2.0)) {
2943 let (_, singular_values, _) = three(backend.svd(&matrix)?, "norm")?;
2944 return if ord == Some(2.0) {
2945 reduce_max(&singular_values, &[0], backend)
2946 } else {
2947 reduce_min(&singular_values, &[0], backend)
2948 };
2949 }
2950
2951 let abs = backend.abs_read(TensorRead::from_tensor(&matrix))?;
2952 match ord {
2953 None => frobenius_norm(&abs, &[0, 1], backend),
2954 Some(p) if p == f64::INFINITY => row_sum_norm(&abs, true, backend),
2955 Some(p) if p == f64::NEG_INFINITY => row_sum_norm(&abs, false, backend),
2956 Some(1.0) => col_sum_norm(&abs, true, backend),
2957 Some(-1.0) => col_sum_norm(&abs, false, backend),
2958 Some(0.0) => count_nonzero(&abs, &[0, 1], backend),
2959 Some(p) => p_norm(&abs, &[0, 1], p, backend),
2960 }
2961}
2962
2963fn row_sum_norm<B: LinalgBackend + ?Sized>(
2964 input: &Tensor,
2965 take_max: bool,
2966 backend: &mut B,
2967) -> tenferro_tensor::Result<Tensor> {
2968 let sums = reduce_sum(input, &[1], backend)?;
2969 if take_max {
2970 reduce_max(&sums, &[0], backend)
2971 } else {
2972 reduce_min(&sums, &[0], backend)
2973 }
2974}
2975
2976fn col_sum_norm<B: LinalgBackend + ?Sized>(
2977 input: &Tensor,
2978 take_max: bool,
2979 backend: &mut B,
2980) -> tenferro_tensor::Result<Tensor> {
2981 let sums = reduce_sum(input, &[0], backend)?;
2982 if take_max {
2983 reduce_max(&sums, &[0], backend)
2984 } else {
2985 reduce_min(&sums, &[0], backend)
2986 }
2987}
2988
2989fn move_axes_to_front<B: LinalgBackend + ?Sized>(
2990 input: TensorRead<'_>,
2991 axes: &[usize],
2992 backend: &mut B,
2993) -> tenferro_tensor::Result<Tensor> {
2994 if axes.iter().enumerate().all(|(index, &axis)| index == axis) {
2995 return backend.to_contiguous_read(input);
2996 }
2997 let mut selected = vec![false; input.shape().len()];
2998 for &axis in axes {
2999 selected[axis] = true;
3000 }
3001 let mut perm = axes.to_vec();
3002 perm.extend(
3003 selected
3004 .iter()
3005 .enumerate()
3006 .filter_map(|(axis, selected)| (!selected).then_some(axis)),
3007 );
3008 backend.transpose_read(input, &perm)
3009}