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}