Skip to main content

tenferro_linalg/
cpu_kernels.rs

1//! Injectable CPU linear-algebra kernels.
2//!
3//! A [`CpuLinalgKernels`] implementation replaces tenferro-linalg's built-in
4//! faer/LAPACK kernels op by op on one CPU backend. Install it with
5//! [`install_linalg_kernels`] on a [`CpuProviderBundleBuilder`]; the CPU
6//! session then asks it first for every primitive it dispatches (Cholesky,
7//! triangular solve, LU and full-pivot LU, solve, SVD, QR and rank-revealing
8//! QR, eigh, eig, and their value-only forms, including the `_read`
9//! variants). A method returns [`CpuLinalgOutcome::Unsupported`], its default,
10//! before producing anything, and the built-in kernel runs instead. The
11//! Householder family, `lu_factor` and prepared LU solves, and `_into`
12//! outputs always use the built-in kernels. Composites (`det`, `slogdet`, `inv`, `lstsq`, `pinv`,
13//! norms) are built from the primitives and follow them.
14//!
15//! A kernel runs inside the session's entered [`CpuExecutionContext`] on host
16//! operands and must return exactly what the built-in kernel returns for the
17//! same input: the same outputs in the same order, shapes, dtypes, pivot and
18//! ordering conventions, and batch handling (trailing batch axes). It may run
19//! in parallel only in [`tenferro_cpu::ParallelMode::Inner`], within
20//! `thread_budget()`, on the context's pool
21//! ([`CpuExecutionContext::rayon_pool`]).
22//!
23//! # Examples
24//!
25//! ```
26//! use std::sync::Arc;
27//! use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
28//! use tenferro_linalg::cpu_kernels::{install_linalg_kernels, CpuLinalgKernels, CpuLinalgKernelsSlot};
29//!
30//! /// Declines everything, so the built-in kernels run.
31//! #[derive(Debug)]
32//! struct Nothing;
33//! impl CpuLinalgKernels for Nothing {}
34//!
35//! let builder = CpuProviderBundle::builder(CpuBackendKind::default_compiled());
36//! let bundle = install_linalg_kernels(builder, Arc::new(Nothing)).build()?;
37//! // Install with `CpuBackend::with_provider_bundle` (a faer backend; explicit
38//! // bundles on the BLAS kind are rejected because BLAS threading is not
39//! // enforceable).
40//! assert!(bundle.extension::<CpuLinalgKernelsSlot>().is_some());
41//! # Ok::<(), Box<dyn std::error::Error>>(())
42//! ```
43
44use std::sync::Arc;
45
46use tenferro_cpu::provider::CpuProviderUnsupported;
47use tenferro_cpu::{CpuExecutionContext, CpuProviderBundleBuilder};
48use tenferro_tensor::{Tensor, TensorView};
49
50use crate::RankRevealingQrOptions;
51
52/// Result of a [`CpuLinalgKernels`] call.
53///
54/// # Examples
55///
56/// ```
57/// use tenferro_cpu::provider::CpuProviderUnsupported;
58/// use tenferro_linalg::cpu_kernels::CpuLinalgOutcome;
59/// let o: CpuLinalgOutcome<u8> = CpuLinalgOutcome::Unsupported(CpuProviderUnsupported::RuntimeUnavailable);
60/// assert!(matches!(o, CpuLinalgOutcome::Unsupported(_)));
61/// ```
62#[derive(Debug)]
63pub enum CpuLinalgOutcome<T> {
64    /// The kernel produced the outputs.
65    Executed(T),
66    /// The kernel declined; nothing was produced and the built-in kernel
67    /// runs.
68    Unsupported(CpuProviderUnsupported),
69}
70
71/// Flags of a triangular solve (`op(A) X = B` or `X op(A) = B`).
72///
73/// # Examples
74///
75/// ```
76/// let o = tenferro_linalg::cpu_kernels::TriangularSolveOptions {
77///     left_side: true, lower: true, transpose_a: false, unit_diagonal: false,
78/// };
79/// assert!(o.lower);
80/// ```
81#[derive(Clone, Copy, Debug, PartialEq, Eq)]
82pub struct TriangularSolveOptions {
83    /// Solve `op(A) X = B` (else `X op(A) = B`).
84    pub left_side: bool,
85    /// `A` is lower triangular.
86    pub lower: bool,
87    /// Use `A^T`.
88    pub transpose_a: bool,
89    /// The diagonal of `A` is taken as ones.
90    pub unit_diagonal: bool,
91}
92
93fn declined<T>() -> tenferro_tensor::Result<CpuLinalgOutcome<T>> {
94    Ok(CpuLinalgOutcome::Unsupported(
95        CpuProviderUnsupported::RuntimeUnavailable,
96    ))
97}
98
99/// Replacement CPU linear-algebra kernels. Every method defaults to
100/// [`CpuLinalgOutcome::Unsupported`]; implement the ones the provider
101/// handles. See the [module documentation](self) for the contract.
102///
103/// # Errors
104///
105/// Methods return errors only for failures after committing to the
106/// operation (a runtime or storage failure); declining is
107/// [`CpuLinalgOutcome::Unsupported`].
108///
109/// # Examples
110///
111/// ```
112/// use tenferro_linalg::cpu_kernels::CpuLinalgKernels;
113/// #[derive(Debug)]
114/// struct Nothing;
115/// impl CpuLinalgKernels for Nothing {}
116/// let _: &dyn CpuLinalgKernels = &Nothing;
117/// ```
118#[allow(unused_variables)]
119pub trait CpuLinalgKernels: std::fmt::Debug + Send + Sync + 'static {
120    /// Cholesky factor, as the built-in `cholesky`.
121    ///
122    /// # Errors
123    ///
124    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
125    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
126    /// host buffer fails after it committed to the operation.
127    fn cholesky(
128        &self,
129        context: &CpuExecutionContext<'_>,
130        input: TensorView<'_>,
131    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
132        declined()
133    }
134
135    /// Triangular solve, as the built-in `triangular_solve`.
136    ///
137    /// # Errors
138    ///
139    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
140    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
141    /// host buffer fails after it committed to the operation.
142    fn triangular_solve(
143        &self,
144        context: &CpuExecutionContext<'_>,
145        a: TensorView<'_>,
146        b: TensorView<'_>,
147        options: TriangularSolveOptions,
148    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
149        declined()
150    }
151
152    /// Partial-pivot LU, as the built-in `lu`.
153    ///
154    /// # Errors
155    ///
156    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
157    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
158    /// host buffer fails after it committed to the operation.
159    fn lu(
160        &self,
161        context: &CpuExecutionContext<'_>,
162        input: TensorView<'_>,
163    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
164        declined()
165    }
166
167    /// Full-pivot LU, as the built-in `full_piv_lu`.
168    ///
169    /// # Errors
170    ///
171    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
172    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
173    /// host buffer fails after it committed to the operation.
174    fn full_piv_lu(
175        &self,
176        context: &CpuExecutionContext<'_>,
177        input: TensorView<'_>,
178    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
179        declined()
180    }
181
182    /// Linear solve `A X = B`, as the built-in `solve`.
183    ///
184    /// # Errors
185    ///
186    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
187    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
188    /// host buffer fails after it committed to the operation.
189    fn solve(
190        &self,
191        context: &CpuExecutionContext<'_>,
192        a: TensorView<'_>,
193        b: TensorView<'_>,
194    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
195        declined()
196    }
197
198    /// Thin SVD, as the built-in `svd`.
199    ///
200    /// # Errors
201    ///
202    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
203    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
204    /// host buffer fails after it committed to the operation.
205    fn svd(
206        &self,
207        context: &CpuExecutionContext<'_>,
208        input: TensorView<'_>,
209    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
210        declined()
211    }
212
213    /// Full SVD, as the built-in `svd_full`.
214    ///
215    /// # Errors
216    ///
217    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
218    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
219    /// host buffer fails after it committed to the operation.
220    fn svd_full(
221        &self,
222        context: &CpuExecutionContext<'_>,
223        input: TensorView<'_>,
224    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
225        declined()
226    }
227
228    /// Singular values, as the built-in `svd_values`.
229    ///
230    /// # Errors
231    ///
232    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
233    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
234    /// host buffer fails after it committed to the operation.
235    fn svd_values(
236        &self,
237        context: &CpuExecutionContext<'_>,
238        input: TensorView<'_>,
239    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
240        declined()
241    }
242
243    /// Thin QR, as the built-in `qr`.
244    ///
245    /// # Errors
246    ///
247    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
248    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
249    /// host buffer fails after it committed to the operation.
250    fn qr(
251        &self,
252        context: &CpuExecutionContext<'_>,
253        input: TensorView<'_>,
254    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
255        declined()
256    }
257
258    /// Column-pivoted QR, as the built-in `rank_revealing_qr`.
259    ///
260    /// # Errors
261    ///
262    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
263    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
264    /// host buffer fails after it committed to the operation.
265    fn rank_revealing_qr(
266        &self,
267        context: &CpuExecutionContext<'_>,
268        input: TensorView<'_>,
269        options: RankRevealingQrOptions,
270    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
271        declined()
272    }
273
274    /// Hermitian eigendecomposition, as the built-in `eigh`.
275    ///
276    /// # Errors
277    ///
278    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
279    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
280    /// host buffer fails after it committed to the operation.
281    fn eigh(
282        &self,
283        context: &CpuExecutionContext<'_>,
284        input: TensorView<'_>,
285    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
286        declined()
287    }
288
289    /// Hermitian eigenvalues, as the built-in `eigh_values`.
290    ///
291    /// # Errors
292    ///
293    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
294    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
295    /// host buffer fails after it committed to the operation.
296    fn eigh_values(
297        &self,
298        context: &CpuExecutionContext<'_>,
299        input: TensorView<'_>,
300    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
301        declined()
302    }
303
304    /// General eigendecomposition, as the built-in `eig`.
305    ///
306    /// # Errors
307    ///
308    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
309    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
310    /// host buffer fails after it committed to the operation.
311    fn eig(
312        &self,
313        context: &CpuExecutionContext<'_>,
314        input: TensorView<'_>,
315    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Vec<Tensor>>> {
316        declined()
317    }
318
319    /// General eigenvalues, as the built-in `eig_values`.
320    ///
321    /// # Errors
322    ///
323    /// Returns [`tenferro_tensor::Error::BackendFailure`] or
324    /// [`tenferro_tensor::Error::BackendSource`] when the kernel's runtime or a
325    /// host buffer fails after it committed to the operation.
326    fn eig_values(
327        &self,
328        context: &CpuExecutionContext<'_>,
329        input: TensorView<'_>,
330    ) -> tenferro_tensor::Result<CpuLinalgOutcome<Tensor>> {
331        declined()
332    }
333}
334
335/// The provider-bundle extension that carries a [`CpuLinalgKernels`].
336///
337/// # Examples
338///
339/// ```
340/// use std::sync::Arc;
341/// use tenferro_linalg::cpu_kernels::{CpuLinalgKernels, CpuLinalgKernelsSlot};
342/// #[derive(Debug)]
343/// struct Nothing;
344/// impl CpuLinalgKernels for Nothing {}
345/// let slot = CpuLinalgKernelsSlot(Arc::new(Nothing));
346/// let _ = &slot.0;
347/// ```
348#[derive(Debug)]
349pub struct CpuLinalgKernelsSlot(pub Arc<dyn CpuLinalgKernels>);
350
351/// Install `kernels` on a provider bundle under construction.
352///
353/// # Examples
354///
355/// ```
356/// use std::sync::Arc;
357/// use tenferro_cpu::{CpuBackendKind, CpuProviderBundle};
358/// use tenferro_linalg::cpu_kernels::{install_linalg_kernels, CpuLinalgKernels, CpuLinalgKernelsSlot};
359/// #[derive(Debug)]
360/// struct Nothing;
361/// impl CpuLinalgKernels for Nothing {}
362/// let bundle = install_linalg_kernels(
363///     CpuProviderBundle::builder(CpuBackendKind::default_compiled()),
364///     Arc::new(Nothing),
365/// )
366/// .build()?;
367/// assert!(bundle.extension::<CpuLinalgKernelsSlot>().is_some());
368/// # Ok::<(), tenferro_cpu::CpuProviderBundleBuildError>(())
369/// ```
370pub fn install_linalg_kernels(
371    builder: CpuProviderBundleBuilder,
372    kernels: Arc<dyn CpuLinalgKernels>,
373) -> CpuProviderBundleBuilder {
374    builder.extension(Arc::new(CpuLinalgKernelsSlot(kernels)))
375}