use core::fmt;
use std::mem::MaybeUninit;
use std::num::NonZeroUsize;
use std::sync::Arc;
use tenferro_tensor::backend::GroupedGemmJob;
use tenferro_tensor::{
DType, DotGeneralAccumulation, Tensor, TensorRead, TensorView, TensorViewMut, TensorWrite,
};
use crate::arbiter::{with_execution_owner, ResourcePermit};
use crate::backend::CpuBackendKind;
use crate::buffer_pool::BufferPool;
use crate::domain_executor::{indexed_jobs, install_scoped};
#[cfg(feature = "cpu-blas")]
use crate::provider_capability::builtin_blas_execution_capabilities;
#[cfg(not(feature = "cpu-blas"))]
use crate::provider_capability::serial_capabilities;
use crate::provider_capability::{engine_worker_capabilities, CpuProviderExecutionCapabilities};
use crate::resource_domain::CpuResourceDomain;
use crate::{
CpuDomainExecutorError, CpuDomainId, CpuInnerParallelism, CpuPlacementGuarantee, CpuSet,
};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CpuOperand {
Lhs,
Rhs,
Output,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum CpuProviderUnsupported {
DType(DType),
Rank {
lhs: usize,
rhs: usize,
},
Layout(CpuOperand),
Conjugation,
Accumulation,
StridedBatch,
Grouped,
RuntimeUnavailable,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[must_use]
pub enum CpuProviderOutcome {
Executed,
Unsupported(CpuProviderUnsupported),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ParallelMode {
Sequential,
Outer,
Inner,
}
#[derive(Clone, Copy)]
pub struct CpuExecutionContext<'a> {
domain: &'a CpuResourceDomain,
parallel_mode: ParallelMode,
}
impl fmt::Debug for CpuExecutionContext<'_> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CpuExecutionContext")
.field("domain_id", &self.domain_id())
.field("cpus", &self.cpus())
.field("thread_budget", &self.thread_budget())
.field("placement_guarantee", &self.placement_guarantee())
.field("parallel_mode", &self.parallel_mode())
.finish_non_exhaustive()
}
}
impl<'a> CpuExecutionContext<'a> {
fn entered(domain: &'a CpuResourceDomain, parallel_mode: ParallelMode) -> Self {
Self {
domain,
parallel_mode,
}
}
pub fn domain_id(&self) -> CpuDomainId {
self.domain.id()
}
pub fn cpus(&self) -> Option<&CpuSet> {
self.domain.cpus()
}
pub fn admission_mode(&self) -> crate::CpuAdmissionMode {
self.domain.admission_mode()
}
pub fn thread_budget(&self) -> NonZeroUsize {
self.domain.thread_budget()
}
pub fn placement_guarantee(&self) -> Option<CpuPlacementGuarantee> {
self.domain.placement_guarantee()
}
pub fn parallel_mode(&self) -> ParallelMode {
self.parallel_mode
}
#[doc(hidden)]
pub fn with_materialized_tensor_read<R>(
&self,
buffers: &mut BufferPool,
op: &'static str,
input: TensorRead<'_>,
operation: impl FnOnce(&Tensor, &mut BufferPool) -> tenferro_tensor::Result<R>,
) -> tenferro_tensor::Result<R> {
match input {
TensorRead::Tensor(tensor) => operation(tensor, buffers),
TensorRead::View(view) => {
let materialized = self.with_native_parallelism(|| {
crate::materialize_tensor_read(buffers, op, TensorRead::View(view))
})?;
let result = operation(&materialized, buffers);
reclaim_tensor(buffers, materialized);
result
}
}
}
#[doc(hidden)]
pub fn reshape_tensor(
&self,
input: &Tensor,
shape: &[usize],
) -> tenferro_tensor::Result<Tensor> {
crate::structural::reshape(input, shape)
}
#[cfg(feature = "cpu-faer")]
#[doc(hidden)]
pub fn faer_parallelism(self) -> faer::Par {
match (
self.parallel_mode,
self.domain.executor_capabilities().inner_parallelism,
) {
(ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
faer::Par::rayon(self.thread_budget().get())
}
_ => faer::Par::Seq,
}
}
pub(crate) fn strided_exec_context(&self) -> strided_kernel::ExecContext {
match (
self.parallel_mode,
self.domain.executor_capabilities().inner_parallelism,
) {
(ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
match strided_kernel::ExecContext::max_threads(self.thread_budget().get()) {
Ok(context) => context,
Err(_) => unreachable!("CpuExecutionContext has a non-zero thread budget"),
}
}
_ => strided_kernel::ExecContext::serial(),
}
}
pub(crate) fn with_native_parallelism<R>(&self, operation: impl FnOnce() -> R) -> R {
let policy = match (
self.parallel_mode,
self.domain.executor_capabilities().inner_parallelism,
) {
(ParallelMode::Inner, CpuInnerParallelism::Rayon) if self.thread_budget().get() > 1 => {
strided_kernel::ExecutionPolicy::Rayon {
max_threads: self.thread_budget(),
}
}
_ => strided_kernel::ExecutionPolicy::Sequential,
};
strided_kernel::with_execution_policy(policy, operation)
}
}
fn reclaim_tensor(buffers: &mut BufferPool, tensor: Tensor) {
match tensor {
Tensor::F32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::F64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::I32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::I64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::Bool(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::C32(tensor) => crate::backend::reclaim_typed(buffers, tensor),
Tensor::C64(tensor) => crate::backend::reclaim_typed(buffers, tensor),
}
}
#[derive(Clone, Copy)]
pub(crate) struct CpuOperationEntry<'a> {
domain: &'a CpuResourceDomain,
permit: &'a ResourcePermit,
}
impl<'a> CpuOperationEntry<'a> {
pub(crate) fn new(domain: &'a CpuResourceDomain, permit: &'a ResourcePermit) -> Self {
Self { domain, permit }
}
pub(crate) fn domain_id(self) -> CpuDomainId {
self.domain.id()
}
pub(crate) fn enter<R: Send>(
self,
parallel_mode: ParallelMode,
operation: impl FnOnce(&CpuExecutionContext<'_>) -> R + Send,
) -> Result<R, CpuDomainExecutorError> {
if parallel_mode == ParallelMode::Outer {
return Err(CpuDomainExecutorError::Scheduling {
message: "CPU executor install requires Sequential or Inner mode, got Outer"
.to_owned(),
});
}
let owner = self.permit.owner();
with_execution_owner(owner, || {
install_scoped(self.domain.executor().as_ref(), || {
with_execution_owner(owner, || {
let context = CpuExecutionContext::entered(self.domain, parallel_mode);
operation(&context)
})
})
})
}
pub(crate) fn enter_or_reuse<R: Send>(
self,
entered: Option<&CpuExecutionContext<'_>>,
parallel_mode: ParallelMode,
operation: impl FnOnce(&CpuExecutionContext<'_>) -> R + Send,
) -> Result<R, CpuDomainExecutorError> {
let Some(entered) = entered else {
return self.enter(parallel_mode, operation);
};
if parallel_mode == ParallelMode::Outer {
return Err(CpuDomainExecutorError::Scheduling {
message: "entered CPU session requires Sequential or Inner mode, got Outer"
.to_owned(),
});
}
if entered.domain_id() != self.domain.id() {
return Err(CpuDomainExecutorError::Scheduling {
message: format!(
"entered CPU session domain {:?} does not match operation domain {:?}",
entered.domain_id(),
self.domain.id()
),
});
}
let owner = self.permit.owner();
Ok(with_execution_owner(owner, || {
let context = CpuExecutionContext::entered(self.domain, parallel_mode);
operation(&context)
}))
}
pub(crate) fn supports_infallible_session_entry(self) -> bool {
self.domain.ownership() == crate::CpuDomainOwnership::Managed
}
pub(crate) fn enter_managed_session<R: Send>(
self,
operation: impl FnOnce(CpuExecutionContext<'a>) -> R + Send,
) -> R {
assert!(
self.supports_infallible_session_entry(),
"managed session entry requires a Tenferro-managed CPU domain"
);
let mode = self.preferred_engine_mode();
self.enter(mode, |_| {
operation(CpuExecutionContext::entered(self.domain, mode))
})
.unwrap_or_else(|error| {
panic!("Tenferro-managed CPU executor violated synchronous install contract: {error}")
})
}
pub(crate) fn submit_outer(
self,
len: usize,
operation: impl Fn(usize, &CpuExecutionContext<'_>) -> Result<(), CpuDomainExecutorError> + Sync,
) -> Result<(), CpuDomainExecutorError> {
if !self.supports_outer() {
return Err(CpuDomainExecutorError::Scheduling {
message: format!(
"CPU domain {:?} does not support Outer mode",
self.domain.id()
),
});
}
let owner = self.permit.owner();
let lane_count = len.min(self.domain.thread_budget().get());
let jobs = indexed_jobs(lane_count, |lane| {
let mut index = lane;
while index < len {
with_execution_owner(owner, || {
let context =
CpuExecutionContext::entered(self.domain, ParallelMode::Sequential);
operation(index, &context)
})?;
let Some(next) = index.checked_add(lane_count) else {
break;
};
index = next;
}
Ok(())
});
with_execution_owner(owner, || self.domain.executor().submit(&jobs))?;
if let Some(index) = jobs.invalid_index_attempt() {
return Err(CpuDomainExecutorError::Scheduling {
message: format!(
"executor requested scoped CPU lane index {index}, but the submission has {lane_count} lanes for {len} logical jobs"
),
});
}
Ok(())
}
pub(crate) fn preferred_engine_mode(self) -> ParallelMode {
if self.domain.thread_budget().get() > 1
&& self.domain.executor_capabilities().inner_parallelism == CpuInnerParallelism::Rayon
{
ParallelMode::Inner
} else {
ParallelMode::Sequential
}
}
pub(crate) fn preferred_provider_mode(
self,
accepts: impl Fn(ParallelMode) -> bool,
) -> Result<ParallelMode, crate::CpuProviderDomainError> {
if self.domain.thread_budget().get() == 1 {
return if accepts(ParallelMode::Sequential) {
Ok(ParallelMode::Sequential)
} else {
Err(crate::CpuProviderDomainError::ParallelModeNotSupported {
mode: ParallelMode::Sequential,
})
};
}
if accepts(ParallelMode::Inner) {
return Ok(ParallelMode::Inner);
}
if accepts(ParallelMode::Sequential) {
return Ok(ParallelMode::Sequential);
}
Err(crate::CpuProviderDomainError::ParallelModeNotSupported {
mode: ParallelMode::Inner,
})
}
pub(crate) fn preferred_linalg_mode(self, kind: CpuBackendKind) -> ParallelMode {
if self.domain.thread_budget().get() == 1 {
return ParallelMode::Sequential;
}
match kind {
CpuBackendKind::Faer => self.preferred_engine_mode(),
CpuBackendKind::Blas => ParallelMode::Inner,
}
}
pub(crate) fn provider_default_compatibility_mode(self) -> ParallelMode {
if self.domain.thread_budget().get() == 1 {
ParallelMode::Sequential
} else {
ParallelMode::Inner
}
}
pub(crate) fn thread_budget(self) -> NonZeroUsize {
self.domain.thread_budget()
}
pub(crate) fn supports_outer(self) -> bool {
self.domain.thread_budget().get() > 1
&& self.domain.executor_capabilities().outer_parallelism
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CpuBatchedMatrixLayout {
offset: isize,
row_stride: isize,
column_stride: isize,
batch_stride: isize,
}
impl CpuBatchedMatrixLayout {
#[allow(dead_code)]
pub(crate) fn new(
offset: isize,
row_stride: isize,
column_stride: isize,
batch_stride: isize,
) -> Self {
Self {
offset,
row_stride,
column_stride,
batch_stride,
}
}
pub fn offset(self) -> isize {
self.offset
}
pub fn row_stride(self) -> isize {
self.row_stride
}
pub fn column_stride(self) -> isize {
self.column_stride
}
pub fn batch_stride(self) -> isize {
self.batch_stride
}
}
#[derive(Debug)]
pub struct CpuGemmRequest<'request, 'input, 'output> {
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
rows: usize,
columns: usize,
contracted: usize,
batch_count: usize,
lhs_layout: CpuBatchedMatrixLayout,
rhs_layout: CpuBatchedMatrixLayout,
output_layout: CpuBatchedMatrixLayout,
accumulation: DotGeneralAccumulation,
}
pub(crate) struct CpuGemmRequestParts<'request, 'input, 'output> {
pub(crate) lhs: &'request TensorRead<'input>,
pub(crate) rhs: &'request TensorRead<'input>,
pub(crate) output: &'request mut TensorWrite<'output>,
pub(crate) rows: usize,
pub(crate) columns: usize,
pub(crate) contracted: usize,
pub(crate) batch_count: usize,
pub(crate) lhs_layout: CpuBatchedMatrixLayout,
pub(crate) rhs_layout: CpuBatchedMatrixLayout,
pub(crate) output_layout: CpuBatchedMatrixLayout,
pub(crate) accumulation: DotGeneralAccumulation,
}
impl<'request, 'input, 'output> CpuGemmRequest<'request, 'input, 'output> {
#[allow(clippy::too_many_arguments, dead_code)]
pub(crate) fn new(
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
rows: usize,
columns: usize,
contracted: usize,
batch_count: usize,
lhs_layout: CpuBatchedMatrixLayout,
rhs_layout: CpuBatchedMatrixLayout,
output_layout: CpuBatchedMatrixLayout,
accumulation: DotGeneralAccumulation,
) -> Self {
Self {
lhs,
rhs,
output,
rows,
columns,
contracted,
batch_count,
lhs_layout,
rhs_layout,
output_layout,
accumulation,
}
}
pub fn lhs(&self) -> &TensorRead<'input> {
self.lhs
}
pub fn rhs(&self) -> &TensorRead<'input> {
self.rhs
}
pub fn output(&mut self) -> &mut TensorWrite<'output> {
self.output
}
pub fn rows(&self) -> usize {
self.rows
}
pub fn columns(&self) -> usize {
self.columns
}
pub fn contracted(&self) -> usize {
self.contracted
}
pub fn batch_count(&self) -> usize {
self.batch_count
}
pub fn lhs_layout(&self) -> CpuBatchedMatrixLayout {
self.lhs_layout
}
pub fn rhs_layout(&self) -> CpuBatchedMatrixLayout {
self.rhs_layout
}
pub fn output_layout(&self) -> CpuBatchedMatrixLayout {
self.output_layout
}
pub fn accumulation(&self) -> DotGeneralAccumulation {
self.accumulation
}
pub(crate) fn into_parts(self) -> CpuGemmRequestParts<'request, 'input, 'output> {
CpuGemmRequestParts {
lhs: self.lhs,
rhs: self.rhs,
output: self.output,
rows: self.rows,
columns: self.columns,
contracted: self.contracted,
batch_count: self.batch_count,
lhs_layout: self.lhs_layout,
rhs_layout: self.rhs_layout,
output_layout: self.output_layout,
accumulation: self.accumulation,
}
}
}
#[derive(Debug)]
pub struct CpuGemmUninitRequest<'request, 'input> {
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
rows: usize,
columns: usize,
contracted: usize,
batch_count: usize,
lhs_layout: CpuBatchedMatrixLayout,
rhs_layout: CpuBatchedMatrixLayout,
output_layout: CpuBatchedMatrixLayout,
accumulation: DotGeneralAccumulation,
}
#[cfg(feature = "cpu-faer")]
pub(crate) struct CpuGemmUninitRequestParts<'request, 'input> {
pub(crate) lhs: &'request TensorRead<'input>,
pub(crate) rhs: &'request TensorRead<'input>,
pub(crate) rows: usize,
pub(crate) columns: usize,
pub(crate) contracted: usize,
pub(crate) batch_count: usize,
pub(crate) lhs_layout: CpuBatchedMatrixLayout,
pub(crate) rhs_layout: CpuBatchedMatrixLayout,
pub(crate) output_layout: CpuBatchedMatrixLayout,
pub(crate) accumulation: DotGeneralAccumulation,
}
impl<'request, 'input> CpuGemmUninitRequest<'request, 'input> {
#[allow(clippy::too_many_arguments, dead_code)]
pub(crate) fn new(
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
rows: usize,
columns: usize,
contracted: usize,
batch_count: usize,
lhs_layout: CpuBatchedMatrixLayout,
rhs_layout: CpuBatchedMatrixLayout,
output_layout: CpuBatchedMatrixLayout,
accumulation: DotGeneralAccumulation,
) -> Self {
Self {
lhs,
rhs,
rows,
columns,
contracted,
batch_count,
lhs_layout,
rhs_layout,
output_layout,
accumulation,
}
}
pub fn lhs(&self) -> &TensorRead<'input> {
self.lhs
}
pub fn rhs(&self) -> &TensorRead<'input> {
self.rhs
}
pub fn rows(&self) -> usize {
self.rows
}
pub fn columns(&self) -> usize {
self.columns
}
pub fn contracted(&self) -> usize {
self.contracted
}
pub fn batch_count(&self) -> usize {
self.batch_count
}
pub fn lhs_layout(&self) -> CpuBatchedMatrixLayout {
self.lhs_layout
}
pub fn rhs_layout(&self) -> CpuBatchedMatrixLayout {
self.rhs_layout
}
pub fn output_layout(&self) -> CpuBatchedMatrixLayout {
self.output_layout
}
pub fn accumulation(&self) -> DotGeneralAccumulation {
self.accumulation
}
#[cfg(feature = "cpu-faer")]
pub(crate) fn into_parts(self) -> CpuGemmUninitRequestParts<'request, 'input> {
CpuGemmUninitRequestParts {
lhs: self.lhs,
rhs: self.rhs,
rows: self.rows,
columns: self.columns,
contracted: self.contracted,
batch_count: self.batch_count,
lhs_layout: self.lhs_layout,
rhs_layout: self.rhs_layout,
output_layout: self.output_layout,
accumulation: self.accumulation,
}
}
}
#[derive(Debug)]
pub struct CpuGroupedGemmRequest<'request, 'input, 'output> {
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
jobs: &'request [GroupedGemmJob],
accumulation: DotGeneralAccumulation,
}
impl<'request, 'input, 'output> CpuGroupedGemmRequest<'request, 'input, 'output> {
#[allow(dead_code)]
pub(crate) fn new(
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
jobs: &'request [GroupedGemmJob],
accumulation: DotGeneralAccumulation,
) -> Self {
Self {
lhs,
rhs,
output,
jobs,
accumulation,
}
}
pub fn lhs(&self) -> &TensorRead<'input> {
self.lhs
}
pub fn rhs(&self) -> &TensorRead<'input> {
self.rhs
}
pub fn output(&mut self) -> &mut TensorWrite<'output> {
self.output
}
pub fn jobs(&self) -> &[GroupedGemmJob] {
self.jobs
}
pub fn accumulation(&self) -> DotGeneralAccumulation {
self.accumulation
}
pub(crate) fn into_parts(
self,
) -> (
&'request TensorRead<'input>,
&'request TensorRead<'input>,
&'request mut TensorWrite<'output>,
&'request [GroupedGemmJob],
DotGeneralAccumulation,
) {
(
self.lhs,
self.rhs,
self.output,
self.jobs,
self.accumulation,
)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CpuLayoutTransformIntent {
CanonicalColumnMajor,
}
#[derive(Debug)]
pub struct CpuLayoutTransformRequest<'request, 'input, 'output> {
input: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
intent: CpuLayoutTransformIntent,
conjugate: bool,
}
impl<'request, 'input, 'output> CpuLayoutTransformRequest<'request, 'input, 'output> {
#[allow(dead_code)]
pub(crate) fn new(
input: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
intent: CpuLayoutTransformIntent,
conjugate: bool,
) -> Self {
Self {
input,
output,
intent,
conjugate,
}
}
pub fn input(&self) -> &TensorRead<'input> {
self.input
}
pub fn output(&mut self) -> &mut TensorWrite<'output> {
self.output
}
pub fn intent(&self) -> CpuLayoutTransformIntent {
self.intent
}
pub fn conjugate(&self) -> bool {
self.conjugate
}
pub(crate) fn into_parts(
self,
) -> (
&'request TensorRead<'input>,
&'request mut TensorWrite<'output>,
CpuLayoutTransformIntent,
bool,
) {
(self.input, self.output, self.intent, self.conjugate)
}
}
#[derive(Clone, Copy, Debug)]
pub struct CpuContractionAxes<'a> {
lhs_rank: usize,
rhs_rank: usize,
lhs_contracting: &'a [usize],
rhs_contracting: &'a [usize],
lhs_batch: &'a [usize],
rhs_batch: &'a [usize],
lhs_role_mask: Option<u64>,
rhs_role_mask: Option<u64>,
}
impl<'a> CpuContractionAxes<'a> {
#[allow(clippy::too_many_arguments, dead_code)]
pub(crate) fn new(
lhs_rank: usize,
rhs_rank: usize,
lhs_contracting: &'a [usize],
rhs_contracting: &'a [usize],
lhs_batch: &'a [usize],
rhs_batch: &'a [usize],
lhs_role_mask: Option<u64>,
rhs_role_mask: Option<u64>,
) -> Self {
Self {
lhs_rank,
rhs_rank,
lhs_contracting,
rhs_contracting,
lhs_batch,
rhs_batch,
lhs_role_mask,
rhs_role_mask,
}
}
pub fn contracting_pairs(&self) -> impl ExactSizeIterator<Item = (usize, usize)> + '_ {
self.lhs_contracting
.iter()
.copied()
.zip(self.rhs_contracting.iter().copied())
}
pub fn batch_pairs(&self) -> impl ExactSizeIterator<Item = (usize, usize)> + '_ {
self.lhs_batch
.iter()
.copied()
.zip(self.rhs_batch.iter().copied())
}
pub fn lhs_free_axes(&self) -> impl Iterator<Item = usize> + '_ {
(0..self.lhs_rank).filter(move |&axis| !self.lhs_axis_has_role(axis))
}
pub fn rhs_free_axes(&self) -> impl Iterator<Item = usize> + '_ {
(0..self.rhs_rank).filter(move |&axis| !self.rhs_axis_has_role(axis))
}
fn lhs_axis_has_role(&self, axis: usize) -> bool {
self.lhs_role_mask.map_or_else(
|| self.lhs_contracting.contains(&axis) || self.lhs_batch.contains(&axis),
|mask| mask & (1_u64 << axis) != 0,
)
}
fn rhs_axis_has_role(&self, axis: usize) -> bool {
self.rhs_role_mask.map_or_else(
|| self.rhs_contracting.contains(&axis) || self.rhs_batch.contains(&axis),
|mask| mask & (1_u64 << axis) != 0,
)
}
}
#[derive(Debug)]
pub struct CpuDotGeneralRequest<'request, 'input, 'output> {
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
axes: CpuContractionAxes<'request>,
accumulation: DotGeneralAccumulation,
}
impl<'request, 'input, 'output> CpuDotGeneralRequest<'request, 'input, 'output> {
#[allow(dead_code)]
pub(crate) fn new(
lhs: &'request TensorRead<'input>,
rhs: &'request TensorRead<'input>,
output: &'request mut TensorWrite<'output>,
axes: CpuContractionAxes<'request>,
accumulation: DotGeneralAccumulation,
) -> Self {
Self {
lhs,
rhs,
output,
axes,
accumulation,
}
}
pub fn lhs(&self) -> &TensorRead<'input> {
self.lhs
}
pub fn rhs(&self) -> &TensorRead<'input> {
self.rhs
}
pub fn output(&mut self) -> &mut TensorWrite<'output> {
self.output
}
pub fn axes(&self) -> &CpuContractionAxes<'request> {
&self.axes
}
pub fn accumulation(&self) -> DotGeneralAccumulation {
self.accumulation
}
pub fn into_parts(
self,
) -> (
&'request TensorRead<'input>,
&'request TensorRead<'input>,
&'request mut TensorWrite<'output>,
CpuContractionAxes<'request>,
DotGeneralAccumulation,
) {
(
self.lhs,
self.rhs,
self.output,
self.axes,
self.accumulation,
)
}
}
pub trait CpuGemmProvider: fmt::Debug + Send + Sync + 'static {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
fn gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome>;
fn strided_batched_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome>;
fn grouped_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGroupedGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome>;
fn uninit_provider(&self) -> Option<&dyn CpuUninitGemmProvider> {
None
}
}
pub unsafe trait CpuUninitGemmProvider: CpuGemmProvider {
unsafe fn gemm_into_uninit(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmUninitRequest<'_, '_>,
output_bytes: &mut [MaybeUninit<u8>],
) -> tenferro_tensor::Result<CpuProviderOutcome>;
}
pub trait CpuLayoutTransformProvider: fmt::Debug + Send + Sync + 'static {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
fn materialize(
&self,
context: &CpuExecutionContext<'_>,
request: CpuLayoutTransformRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome>;
fn uninit_provider(&self) -> Option<&dyn CpuUninitLayoutTransformProvider> {
None
}
}
pub unsafe trait CpuUninitLayoutTransformProvider: CpuLayoutTransformProvider {
unsafe fn materialize_into_uninit(
&self,
context: &CpuExecutionContext<'_>,
input: &TensorRead<'_>,
intent: CpuLayoutTransformIntent,
conjugate: bool,
output_bytes: &mut [MaybeUninit<u8>],
) -> tenferro_tensor::Result<CpuProviderOutcome>;
}
pub trait CpuGeneralContractionProvider: fmt::Debug + Send + Sync + 'static {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities;
fn dot_general(
&self,
context: &CpuExecutionContext<'_>,
request: CpuDotGeneralRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome>;
}
#[derive(Clone, Copy, Debug, Default)]
pub struct FaerGemmProvider;
impl CpuGemmProvider for FaerGemmProvider {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
engine_worker_capabilities()
}
fn gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
#[cfg(feature = "cpu-faer")]
{
crate::gemm::execute_faer_gemm_request(context, request)
}
#[cfg(not(feature = "cpu-faer"))]
{
let _ = (context, request);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::RuntimeUnavailable,
))
}
}
fn strided_batched_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
self.gemm(context, request)
}
fn grouped_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGroupedGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
#[cfg(feature = "cpu-faer")]
{
crate::gemm::execute_faer_grouped_request(context, request)
}
#[cfg(not(feature = "cpu-faer"))]
{
let _ = (context, request);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::RuntimeUnavailable,
))
}
}
fn uninit_provider(&self) -> Option<&dyn CpuUninitGemmProvider> {
Some(self)
}
}
unsafe impl CpuUninitGemmProvider for FaerGemmProvider {
unsafe fn gemm_into_uninit(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmUninitRequest<'_, '_>,
output_bytes: &mut [MaybeUninit<u8>],
) -> tenferro_tensor::Result<CpuProviderOutcome> {
#[cfg(feature = "cpu-faer")]
{
crate::gemm::execute_faer_gemm_request_into_uninit(context, request, output_bytes)
}
#[cfg(not(feature = "cpu-faer"))]
{
let _ = (context, request, output_bytes);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::RuntimeUnavailable,
))
}
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct BlasGemmProvider;
impl CpuGemmProvider for BlasGemmProvider {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
#[cfg(feature = "cpu-blas")]
{
builtin_blas_execution_capabilities()
}
#[cfg(not(feature = "cpu-blas"))]
{
serial_capabilities()
}
}
fn gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
#[cfg(feature = "cpu-blas")]
{
crate::gemm::execute_blas_gemm_request(context, request)
}
#[cfg(not(feature = "cpu-blas"))]
{
let _ = (context, request);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::RuntimeUnavailable,
))
}
}
fn strided_batched_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
self.gemm(context, request)
}
fn grouped_gemm(
&self,
context: &CpuExecutionContext<'_>,
request: CpuGroupedGemmRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
#[cfg(feature = "cpu-blas")]
{
crate::gemm::execute_blas_grouped_request(context, request)
}
#[cfg(not(feature = "cpu-blas"))]
{
let _ = (context, request);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::RuntimeUnavailable,
))
}
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct StridedLayoutTransformProvider;
impl CpuLayoutTransformProvider for StridedLayoutTransformProvider {
fn execution_capabilities(&self) -> CpuProviderExecutionCapabilities {
engine_worker_capabilities()
}
fn materialize(
&self,
context: &CpuExecutionContext<'_>,
request: CpuLayoutTransformRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
context.with_native_parallelism(|| materialize_strided_layout(request))
}
fn uninit_provider(&self) -> Option<&dyn CpuUninitLayoutTransformProvider> {
Some(self)
}
}
unsafe impl CpuUninitLayoutTransformProvider for StridedLayoutTransformProvider {
unsafe fn materialize_into_uninit(
&self,
context: &CpuExecutionContext<'_>,
input: &TensorRead<'_>,
intent: CpuLayoutTransformIntent,
conjugate: bool,
output_bytes: &mut [MaybeUninit<u8>],
) -> tenferro_tensor::Result<CpuProviderOutcome> {
context.with_native_parallelism(|| {
materialize_strided_layout_into_uninit(input, intent, conjugate, output_bytes)
})
}
}
fn materialize_strided_layout(
request: CpuLayoutTransformRequest<'_, '_, '_>,
) -> tenferro_tensor::Result<CpuProviderOutcome> {
let (input, output, _intent, conjugate) = request.into_parts();
if conjugate {
macro_rules! dispatch_conjugated {
($owned:ident, $view:ident) => {
match (input, &mut *output) {
(
TensorRead::Tensor(Tensor::$owned(input)),
TensorWrite::Tensor(Tensor::$owned(output)),
) => {
let input = input.as_view();
let mut output = output.as_view_mut();
crate::structural::typed_conjugate_view_into(
&input,
&mut output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::View(TensorView::$view(input)),
TensorWrite::Tensor(Tensor::$owned(output)),
) => {
let mut output = output.as_view_mut();
crate::structural::typed_conjugate_view_into(
input,
&mut output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::Tensor(Tensor::$owned(input)),
TensorWrite::View(TensorViewMut::$view(output)),
) => {
let input = input.as_view();
crate::structural::typed_conjugate_view_into(
&input,
output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::View(TensorView::$view(input)),
TensorWrite::View(TensorViewMut::$view(output)),
) => {
crate::structural::typed_conjugate_view_into(
input,
output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
_ => {}
}
};
}
dispatch_conjugated!(F32, F32);
dispatch_conjugated!(F64, F64);
dispatch_conjugated!(C32, C32);
dispatch_conjugated!(C64, C64);
return Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::DType(input.dtype()),
));
}
macro_rules! dispatch {
($owned:ident, $view:ident) => {
match (input, &mut *output) {
(
TensorRead::Tensor(Tensor::$owned(input)),
TensorWrite::Tensor(Tensor::$owned(output)),
) => {
let input = input.as_view();
let mut output = output.as_view_mut();
crate::structural::typed_copy_view_into(
&input,
&mut output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::View(TensorView::$view(input)),
TensorWrite::Tensor(Tensor::$owned(output)),
) => {
let mut output = output.as_view_mut();
crate::structural::typed_copy_view_into(
input,
&mut output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::Tensor(Tensor::$owned(input)),
TensorWrite::View(TensorViewMut::$view(output)),
) => {
let input = input.as_view();
crate::structural::typed_copy_view_into(
&input,
output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
(
TensorRead::View(TensorView::$view(input)),
TensorWrite::View(TensorViewMut::$view(output)),
) => {
crate::structural::typed_copy_view_into(
input,
output,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
_ => {}
}
};
}
dispatch!(F32, F32);
dispatch!(F64, F64);
dispatch!(I32, I32);
dispatch!(I64, I64);
dispatch!(Bool, Bool);
dispatch!(C32, C32);
dispatch!(C64, C64);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::DType(input.dtype()),
))
}
fn materialize_strided_layout_into_uninit(
input: &TensorRead<'_>,
intent: CpuLayoutTransformIntent,
conjugate: bool,
output_bytes: &mut [MaybeUninit<u8>],
) -> tenferro_tensor::Result<CpuProviderOutcome> {
if intent != CpuLayoutTransformIntent::CanonicalColumnMajor {
return Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::Layout(CpuOperand::Output),
));
}
macro_rules! dispatch {
($owned:ident, $view:ident) => {
match input {
TensorRead::Tensor(Tensor::$owned(input)) => {
let input = input.as_view();
crate::structural::typed_copy_into_uninit(
&input,
conjugate,
output_bytes,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
TensorRead::View(TensorView::$view(input)) => {
crate::structural::typed_copy_into_uninit(
input,
conjugate,
output_bytes,
"cpu layout materialization",
)?;
return Ok(CpuProviderOutcome::Executed);
}
_ => {}
}
};
}
dispatch!(F32, F32);
dispatch!(F64, F64);
dispatch!(C32, C32);
dispatch!(C64, C64);
Ok(CpuProviderOutcome::Unsupported(
CpuProviderUnsupported::DType(input.dtype()),
))
}
pub(crate) fn builtin_gemm_provider(kind: CpuBackendKind) -> Arc<dyn CpuGemmProvider> {
match kind {
CpuBackendKind::Faer => Arc::new(FaerGemmProvider),
CpuBackendKind::Blas => Arc::new(BlasGemmProvider),
}
}
pub(crate) fn builtin_layout_provider() -> Arc<dyn CpuLayoutTransformProvider> {
Arc::new(StridedLayoutTransformProvider)
}
#[cfg(test)]
pub(crate) mod tests;