1#![cfg_attr(
22 all(test, feature = "provider-inject"),
23 allow(dead_code, unused_imports)
24)]
25
26#[cfg(not(any(feature = "cpu-faer", feature = "cpu-blas")))]
27compile_error!("enable at least one CPU backend: cpu-faer or cpu-blas");
28
29#[cfg(all(feature = "provider-inject", not(feature = "cpu-blas")))]
30compile_error!("provider-inject requires cpu-blas");
31
32#[cfg(any(
33 all(feature = "blas-openblas", feature = "blas-accelerate"),
34 all(feature = "blas-openblas", feature = "blas-mkl"),
35 all(feature = "blas-accelerate", feature = "blas-mkl"),
36))]
37compile_error!(
38 "enable at most one explicit BLAS provider feature: blas-openblas, blas-accelerate, or blas-mkl"
39);
40
41#[cfg(all(
42 feature = "provider-inject",
43 any(
44 feature = "blas-openblas",
45 feature = "blas-accelerate",
46 feature = "blas-mkl"
47 )
48))]
49compile_error!("provider-inject cannot be combined with explicit BLAS provider features");
50
51pub mod affinity;
52mod affinity_policy;
53mod analytic;
54mod arbiter;
55pub mod backend;
56mod blas1;
57pub(crate) mod buffer_pool {
58 pub use tenferro_internal_cpu_kernels::buffer_pool::*;
59}
60mod capability;
61pub mod context;
62#[allow(dead_code)]
65mod domain_executor;
66#[allow(dead_code)]
67mod dot_runtime;
68pub(crate) use tenferro_internal_cpu_kernels::elementwise;
69pub(crate) use tenferro_internal_cpu_kernels::elementwise::{
70 erased_raw_strided_ref, erased_raw_strided_uninit_mut,
71};
72pub(crate) use tenferro_internal_cpu_kernels::PooledUninitOutput;
73mod engine;
74mod exec_session;
75mod gemm;
76mod indexed_plan_cache;
77mod indexing;
78#[cfg(feature = "provider-inject")]
79pub mod inject;
80mod placement;
81pub mod provider;
82mod provider_capability;
83mod reduction;
84mod resource_domain;
85mod runtime_adapter;
86mod structural;
87mod topology;
88
89use std::ptr::NonNull;
90#[cfg(test)]
91use strided_kernel::StridedArray;
92use strided_kernel::{col_major_strides as kernel_col_major_strides, StridedView};
93
94use crate::buffer_pool::BufferPool;
95pub(crate) use tenferro_tensor::*;
96
97pub(crate) fn cpu_contraction_unsupported_dtype_message(dtype: DType) -> String {
98 let remedy = matches!(dtype, DType::I32 | DType::I64)
99 .then_some(format!("; convert {dtype:?} to F64 before contraction"));
100 format!(
101 "CPU contraction providers support F32/F64/C32/C64{}",
102 remedy.unwrap_or_default()
103 )
104}
105
106pub(crate) fn erased_raw_strided_mut<'a>(
107 dtype: strided_kernel::KernelDType,
108 data: &'a mut [u8],
109 dims: &'a [usize],
110 strides: &'a [isize],
111 offset: isize,
112) -> strided_kernel::Result<strided_kernel::ErasedRawStridedMut<'a>> {
113 let data_ptr = NonNull::new(data.as_mut_ptr()).unwrap_or_else(NonNull::dangling);
114 unsafe {
117 strided_kernel::ErasedRawStridedMut::from_raw_parts(
118 dtype,
119 data_ptr,
120 data.len(),
121 dims,
122 strides,
123 offset,
124 )
125 }
126}
127
128#[cfg(feature = "provider-src")]
129extern crate blas_src as _;
130#[cfg(feature = "provider-inject")]
131extern crate cblas_inject as _;
132#[cfg(feature = "provider-src")]
133extern crate cblas_src as _;
134#[cfg(feature = "provider-inject")]
135extern crate lapack_inject as _;
136#[cfg(feature = "provider-src")]
137extern crate lapack_src as _;
138
139pub use affinity::{
140 available_parallelism, process_cpu_affinity, process_cpu_affinity_count, CpuAffinityError,
141};
142pub use affinity_policy::{
143 resolve_cpu_affinity, resolve_cpu_affinity_with_override, CpuAffinityInput,
144 CpuAffinityInputError, CpuAffinityPolicy, CpuAffinityResolutionError, CpuAffinitySelection,
145 CpuAffinitySelectionReason,
146};
147pub use backend::{
148 CpuBackend, CpuBackendError, CpuBackendKind, CpuExecutionInfo, CpuExecutionMode,
149 CpuRuntimeIdentity, ExternalCpuDomainRegistryError,
150};
151pub use buffer_pool::BufferPoolStats;
152pub use capability::cpu_capabilities;
153pub use context::{CpuContext, CpuContextError};
154pub use domain_executor::{
155 CpuDomainExecutor, CpuDomainExecutorCapabilities, CpuDomainExecutorError, CpuExecutorAffinity,
156 CpuExecutorReentrancy, CpuExecutorShutdown, CpuInnerParallelism, RayonCpuDomainExecutor,
157 ScopedCpuJob, ScopedCpuJobs,
158};
159pub use dot_runtime::{
160 CpuProviderBundle, CpuProviderBundleBuildError, CpuProviderBundleBuilder,
161 CpuProviderBundleInstallError, CpuProviderSlot, GeneralContractionPolicy,
162};
163#[doc(hidden)]
164pub use exec_session::CpuExecSession;
165pub use indexed_plan_cache::IndexedPlanCacheLimits;
166pub use placement::{
167 CpuEngineConstructionError, CpuPlacement, CpuPlacementError, CpuPlacementGuarantee,
168 ResolvedCpuPlacement,
169};
170pub use provider::{CpuExecutionContext, ParallelMode};
171pub use provider_capability::{
172 CpuPlacementControl, CpuProviderDomainError, CpuProviderExecutionCapabilities,
173 CpuThreadCountControl,
174};
175pub use resource_domain::{
176 CpuAdmissionMode, CpuDomainOwnership, ExternalCpuDomain, ExternalCpuDomainError,
177};
178pub use runtime_adapter::{
179 runtime_engine_id, runtime_engine_registration, runtime_engine_registration_with_id,
180 runtime_hardware_class,
181};
182pub use topology::{
183 discover_cpu_topology, CpuId, CpuNode, CpuSet, CpuSetError, CpuTopology, CpuTopologyError,
184 NumaNodeId,
185};
186
187#[doc(hidden)]
194pub fn with_cpu_exec_session<B, R>(
195 session: &mut B,
196 f: impl for<'a> FnOnce(&'a mut CpuExecSession<'a>) -> R,
197) -> Option<R>
198where
199 B: tenferro_tensor::BackendSession + ?Sized,
200{
201 if session.session_type_id() != std::any::TypeId::of::<exec_session::CpuExecSessionMarker>() {
202 return None;
203 }
204 let data = unsafe { session.session_data_mut() };
205 Some(unsafe { f(&mut *(data.cast::<CpuExecSession<'static>>())) })
211}
212
213#[cfg(feature = "cpu-faer")]
240#[cfg_attr(docsrs, doc(cfg(feature = "cpu-faer")))]
241pub trait FaerParallelismExt {
242 fn with_faer_parallelism(
263 &mut self,
264 callback: impl FnOnce(faer::Par) -> tenferro_tensor::Result<()> + Send,
265 ) -> tenferro_tensor::Result<()>;
266}
267
268#[cfg(feature = "cpu-faer")]
269impl<S> FaerParallelismExt for S
270where
271 S: tenferro_tensor::BackendSession + ?Sized,
272{
273 fn with_faer_parallelism(
274 &mut self,
275 callback: impl FnOnce(faer::Par) -> tenferro_tensor::Result<()> + Send,
276 ) -> tenferro_tensor::Result<()> {
277 with_cpu_exec_session(self, |session| session.with_faer_parallelism(callback))
278 .unwrap_or_else(|| {
279 Err(tenferro_tensor::Error::unsupported(
280 "with_faer_parallelism",
281 "selected session is not a CPU/faer execution session",
282 ))
283 })
284 }
285}
286
287#[cfg(test)]
290pub(crate) use analytic::pow;
291#[cfg(test)]
292macro_rules! test_elementwise_wrapper {
293 ($name:ident($($arg:ident: $ty:ty),*) => $with_pool:ident) => {
294 pub(crate) fn $name($($arg: $ty),*) -> crate::Result<Tensor> {
295 let mut buffers = BufferPool::new();
296 elementwise::$with_pool(&mut buffers, $($arg),*)
297 }
298 };
299}
300#[cfg(test)]
301test_elementwise_wrapper!(abs(input: &Tensor) => abs_with_pool);
302#[cfg(test)]
303test_elementwise_wrapper!(add(lhs: &Tensor, rhs: &Tensor) => add_with_pool);
304#[cfg(test)]
305test_elementwise_wrapper!(clamp(input: &Tensor, lower: &Tensor, upper: &Tensor) => clamp_with_pool);
306#[cfg(test)]
307test_elementwise_wrapper!(compare(lhs: &Tensor, rhs: &Tensor, dir: &CompareDir) => compare_with_pool);
308#[cfg(test)]
309test_elementwise_wrapper!(conj(input: &Tensor) => conj_with_pool);
310#[cfg(test)]
311test_elementwise_wrapper!(div(lhs: &Tensor, rhs: &Tensor) => div_with_pool);
312#[cfg(test)]
313test_elementwise_wrapper!(maximum(lhs: &Tensor, rhs: &Tensor) => maximum_with_pool);
314#[cfg(test)]
315test_elementwise_wrapper!(minimum(lhs: &Tensor, rhs: &Tensor) => minimum_with_pool);
316#[cfg(test)]
317test_elementwise_wrapper!(mul(lhs: &Tensor, rhs: &Tensor) => mul_with_pool);
318#[cfg(test)]
319test_elementwise_wrapper!(neg(input: &Tensor) => neg_with_pool);
320#[cfg(test)]
321test_elementwise_wrapper!(rem(lhs: &Tensor, rhs: &Tensor) => rem_with_pool);
322#[cfg(test)]
323test_elementwise_wrapper!(select(pred: &Tensor, on_true: &Tensor, on_false: &Tensor) => select_with_pool);
324#[cfg(test)]
325test_elementwise_wrapper!(sign(input: &Tensor) => sign_with_pool);
326#[cfg(test)]
327test_elementwise_wrapper!(sub(lhs: &Tensor, rhs: &Tensor) => sub_with_pool);
328#[cfg(test)]
329pub(crate) use indexing::{dynamic_slice, dynamic_update_slice, gather, pad, scatter};
330#[cfg(test)]
331pub(crate) use reduction::{reduce_max, reduce_min, reduce_prod, reduce_sum, reduce_sum_squares};
332#[cfg(test)]
333pub(crate) use structural::{
334 broadcast_in_dim, embed_diagonal, extract_diagonal, reshape, transpose, tril, triu,
335};
336
337#[doc(hidden)]
343pub mod linalg_interop {
344 pub use crate::buffer_pool::{BufferPool, PoolScalar};
345 pub use tenferro_internal_cpu_kernels::PooledUninitOutput;
346}
347
348pub(crate) fn cpu_backend_buffer_error(op: &'static str) -> crate::Error {
349 crate::Error::runtime_state(
350 op,
351 "CPU backend received backend buffer; download to host before CPU execution",
352 )
353}
354
355#[derive(Debug, thiserror::Error)]
356pub(crate) enum CpuNumericalError {
357 #[error("{op} received a negative integer exponent for dtype {dtype:?}")]
358 NegativeIntegerExponent { op: &'static str, dtype: DType },
359}
360
361pub(crate) fn cpu_negative_integer_exponent(op: &'static str, dtype: DType) -> crate::Error {
362 crate::Error::extension(
363 op,
364 "cpu",
365 ErrorKind::NumericalFailure,
366 CpuNumericalError::NegativeIntegerExponent { op, dtype },
367 )
368}
369
370pub(crate) trait ConjElem {
371 fn conj_elem(self) -> Self;
372}
373
374impl ConjElem for f32 {
375 fn conj_elem(self) -> Self {
376 self
377 }
378}
379
380impl ConjElem for f64 {
381 fn conj_elem(self) -> Self {
382 self
383 }
384}
385
386impl ConjElem for num_complex::Complex32 {
387 fn conj_elem(self) -> Self {
388 self.conj()
389 }
390}
391
392impl ConjElem for num_complex::Complex64 {
393 fn conj_elem(self) -> Self {
394 self.conj()
395 }
396}
397
398pub(crate) fn typed_host_data<'a, T: TensorScalar>(
399 op: &'static str,
400 tensor: &'a TypedTensor<T>,
401) -> crate::Result<&'a [T]> {
402 if tensor.backend_buffer().is_some() {
403 return Err(cpu_backend_buffer_error(op));
404 }
405 tensor.host_data()
406}
407
408pub(crate) fn typed_view<'a, T: Copy + TensorScalar>(
409 op: &'static str,
410 tensor: &'a TypedTensor<T>,
411) -> crate::Result<StridedView<'a, T>> {
412 if tensor.backend_buffer().is_some() {
413 return Err(cpu_backend_buffer_error(op));
414 }
415 let data = tensor.host_data()?;
416 let strides = kernel_col_major_strides(tensor.shape());
417 StridedView::new(data, tensor.shape(), &strides, 0)
418 .map_err(|err| crate::Error::backend_source(op, err))
419}
420
421pub(crate) fn typed_view_from_view<'a, T: Copy + 'static, R: TensorRank>(
422 op: &'static str,
423 view: &TypedTensorView<'a, T, R>,
424) -> crate::Result<StridedView<'a, T>> {
425 if view.backend_buffer().is_some() {
426 return Err(cpu_backend_buffer_error(op));
427 }
428 StridedView::new(
429 view.host_storage()?,
430 view.shape(),
431 view.strides(),
432 view.offset(),
433 )
434 .map_err(|err| crate::Error::backend_source(op, err))
435}
436
437pub(crate) fn materialize_tensor_read(
438 buffers: &mut BufferPool,
439 op: &'static str,
440 input: TensorRead<'_>,
441) -> crate::Result<Tensor> {
442 match input {
443 TensorRead::Tensor(tensor) => clone_host_tensor_read(op, tensor),
444 TensorRead::View(view) => materialize_tensor_view(buffers, op, view),
445 }
446}
447
448pub(crate) fn copy_tensor_read_into(
449 op: &'static str,
450 src: TensorRead<'_>,
451 dst: TensorWrite<'_>,
452) -> crate::Result<()> {
453 let src_dtype = src.dtype();
454 let dst_dtype = dst.dtype();
455 macro_rules! copy_source {
456 ($variant:ident, $src:expr) => {{
457 let src = $src;
458 match dst {
459 TensorWrite::Tensor(Tensor::$variant(dst)) => {
460 let mut dst = dst.as_view_mut();
461 structural::typed_copy_view_into(&src, &mut dst, op)
462 }
463 TensorWrite::View(TensorViewMut::$variant(mut dst)) => {
464 structural::typed_copy_view_into(&src, &mut dst, op)
465 }
466 _ => Err(crate::Error::dtype_mismatch(op, src_dtype, dst_dtype)),
467 }
468 }};
469 }
470
471 match src {
472 TensorRead::Tensor(Tensor::F32(src)) => copy_source!(F32, src.as_view()),
473 TensorRead::Tensor(Tensor::F64(src)) => copy_source!(F64, src.as_view()),
474 TensorRead::Tensor(Tensor::I32(src)) => copy_source!(I32, src.as_view()),
475 TensorRead::Tensor(Tensor::I64(src)) => copy_source!(I64, src.as_view()),
476 TensorRead::Tensor(Tensor::Bool(src)) => copy_source!(Bool, src.as_view()),
477 TensorRead::Tensor(Tensor::C32(src)) => copy_source!(C32, src.as_view()),
478 TensorRead::Tensor(Tensor::C64(src)) => copy_source!(C64, src.as_view()),
479 TensorRead::View(TensorView::F32(src)) => copy_source!(F32, src),
480 TensorRead::View(TensorView::F64(src)) => copy_source!(F64, src),
481 TensorRead::View(TensorView::I32(src)) => copy_source!(I32, src),
482 TensorRead::View(TensorView::I64(src)) => copy_source!(I64, src),
483 TensorRead::View(TensorView::Bool(src)) => copy_source!(Bool, src),
484 TensorRead::View(TensorView::C32(src)) => copy_source!(C32, src),
485 TensorRead::View(TensorView::C64(src)) => copy_source!(C64, src),
486 }
487}
488
489fn clone_host_tensor_read(op: &'static str, tensor: &Tensor) -> crate::Result<Tensor> {
490 macro_rules! clone_host {
491 ($variant:ident, $tensor:expr) => {{
492 structural::validate_cpu_host_placement(op, "source", $tensor.placement())?;
493 typed_host_data(op, $tensor)?;
494 $tensor.duplicate().map(Tensor::$variant)
495 }};
496 }
497
498 match tensor {
499 Tensor::F32(tensor) => clone_host!(F32, tensor),
500 Tensor::F64(tensor) => clone_host!(F64, tensor),
501 Tensor::I32(tensor) => clone_host!(I32, tensor),
502 Tensor::I64(tensor) => clone_host!(I64, tensor),
503 Tensor::Bool(tensor) => clone_host!(Bool, tensor),
504 Tensor::C32(tensor) => clone_host!(C32, tensor),
505 Tensor::C64(tensor) => clone_host!(C64, tensor),
506 }
507}
508
509fn materialize_tensor_view(
510 buffers: &mut BufferPool,
511 op: &'static str,
512 view: TensorView<'_>,
513) -> crate::Result<Tensor> {
514 macro_rules! materialize {
515 ($variant:ident, $view:expr) => {{
516 Ok(Tensor::$variant(
517 structural::typed_materialize_view_with_pool(buffers, &$view, op)?,
518 ))
519 }};
520 }
521
522 match view {
523 TensorView::F32(view) => materialize!(F32, view),
524 TensorView::F64(view) => materialize!(F64, view),
525 TensorView::I32(view) => materialize!(I32, view),
526 TensorView::I64(view) => materialize!(I64, view),
527 TensorView::Bool(view) => materialize!(Bool, view),
528 TensorView::C32(view) => materialize!(C32, view),
529 TensorView::C64(view) => materialize!(C64, view),
530 }
531}
532
533#[allow(clippy::uninit_vec)]
539#[cfg(test)]
540pub(crate) unsafe fn typed_array_uninit<T>(shape: &[usize]) -> StridedArray<T> {
541 let total: usize = shape.iter().product();
542 let strides = kernel_col_major_strides(shape);
543 let mut data = Vec::with_capacity(total);
544 unsafe { data.set_len(total) };
546 StridedArray::from_parts(data, shape, &strides, 0).expect("column-major output array")
549}
550
551#[cfg(test)]
552pub(crate) fn tensor_from_array<T: Clone + tenferro_tensor::TensorScalar>(
553 array: StridedArray<T>,
554) -> TypedTensor<T> {
555 TypedTensor::from_vec_col_major(array.dims().to_vec(), array.into_data())
557 .expect("strided array dimensions match owned data length")
558}
559
560pub(crate) fn flat_to_multi(mut flat: usize, shape: &[usize], out: &mut [usize]) {
561 assert_eq!(shape.len(), out.len());
562 for (axis, &dim) in shape.iter().enumerate() {
563 if dim == 0 {
564 out[axis] = 0;
565 } else {
566 out[axis] = flat % dim;
567 flat /= dim;
568 }
569 }
570}
571
572#[cfg(all(test, not(feature = "provider-inject")))]
577mod tests;