use std::num::NonZeroUsize;
use thiserror::Error;
use crate::{CpuPlacementGuarantee, CpuSet, ParallelMode};
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub enum CpuThreadCountControl {
Sequential,
PerCallUpperBound,
BinaryClampToOne,
#[default]
GlobalOrUncontrolled,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub enum CpuPlacementControl {
EngineWorkers,
CallingThread,
ExternalWorkers,
#[default]
None,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct CpuProviderExecutionCapabilities {
pub thread_count: CpuThreadCountControl,
pub placement: CpuPlacementControl,
pub worker_local_sequential: bool,
pub accepts_sequential: bool,
pub accepts_outer: bool,
pub accepts_inner: bool,
}
impl Default for CpuProviderExecutionCapabilities {
fn default() -> Self {
Self {
thread_count: CpuThreadCountControl::GlobalOrUncontrolled,
placement: CpuPlacementControl::None,
worker_local_sequential: false,
accepts_sequential: false,
accepts_outer: false,
accepts_inner: true,
}
}
}
#[derive(Clone, Copy, Debug, Eq, Error, PartialEq)]
pub enum CpuProviderDomainError {
#[error(
"provider thread-count control {control:?} cannot enforce thread budget {thread_budget}"
)]
ThreadCountNotEnforceable {
thread_budget: usize,
control: CpuThreadCountControl,
},
#[error(
"provider placement control {placement:?} cannot enforce {guarantee:?} placement for thread budget {thread_budget}"
)]
PlacementNotEnforceable {
thread_budget: usize,
placement: CpuPlacementControl,
guarantee: CpuPlacementGuarantee,
},
#[error(
"provider placement control {placement:?} can leave the caller-managed executor for thread budget {thread_budget}"
)]
CallerManagedPlacementNotEnforceable {
thread_budget: usize,
placement: CpuPlacementControl,
},
#[error("provider cannot honor requested CPU parallel mode {mode:?}")]
ParallelModeNotSupported {
mode: ParallelMode,
},
}
impl CpuProviderExecutionCapabilities {
pub(crate) fn accepts_mode(self, mode: ParallelMode) -> bool {
match mode {
ParallelMode::Sequential => self.accepts_sequential && self.worker_local_sequential,
ParallelMode::Outer => self.accepts_outer && self.worker_local_sequential,
ParallelMode::Inner => self.accepts_inner,
}
}
}
#[cfg(test)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum OpenBlasParallelism {
Sequential,
Pthread,
OpenMp,
Unknown,
}
#[cfg(test)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct OpenBlasProbe {
pub(crate) parallelism: OpenBlasParallelism,
pub(crate) process_global_set_restore_wired: bool,
}
#[cfg(test)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct AccelerateProbe {
pub(crate) binary_thread_local_control_wired: bool,
}
#[cfg(test)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum CpuProviderProbe {
FaerOrNative,
Mkl { thread_local_setter_wired: bool },
OpenBlas(OpenBlasProbe),
Accelerate(AccelerateProbe),
ArmPlOpenMp,
ArmPlSerial,
NvplSerial,
UnknownBlas,
Injected(Option<CpuProviderExecutionCapabilities>),
}
#[cfg(test)]
pub(crate) fn classify_provider(probe: CpuProviderProbe) -> CpuProviderExecutionCapabilities {
match probe {
CpuProviderProbe::FaerOrNative => engine_worker_capabilities(),
CpuProviderProbe::Mkl {
thread_local_setter_wired: true,
} => controlled_external_capabilities(CpuThreadCountControl::PerCallUpperBound),
CpuProviderProbe::Mkl {
thread_local_setter_wired: false,
}
| CpuProviderProbe::ArmPlOpenMp => uncontrolled_external_capabilities(),
CpuProviderProbe::OpenBlas(probe) => classify_openblas(probe),
CpuProviderProbe::Accelerate(probe) => classify_accelerate(probe),
CpuProviderProbe::ArmPlSerial | CpuProviderProbe::NvplSerial => serial_capabilities(),
CpuProviderProbe::UnknownBlas | CpuProviderProbe::Injected(None) => {
CpuProviderExecutionCapabilities::default()
}
CpuProviderProbe::Injected(Some(capabilities)) => capabilities,
}
}
#[cfg(test)]
fn classify_openblas(probe: OpenBlasProbe) -> CpuProviderExecutionCapabilities {
match (probe.parallelism, probe.process_global_set_restore_wired) {
(OpenBlasParallelism::Sequential, _) => serial_capabilities(),
(OpenBlasParallelism::Pthread | OpenBlasParallelism::OpenMp, _) => {
uncontrolled_external_capabilities()
}
(OpenBlasParallelism::Unknown, _) => CpuProviderExecutionCapabilities::default(),
}
}
#[cfg(test)]
fn classify_accelerate(probe: AccelerateProbe) -> CpuProviderExecutionCapabilities {
if probe.binary_thread_local_control_wired {
controlled_external_capabilities(CpuThreadCountControl::BinaryClampToOne)
} else {
uncontrolled_external_capabilities()
}
}
pub(crate) fn engine_worker_capabilities() -> CpuProviderExecutionCapabilities {
CpuProviderExecutionCapabilities {
thread_count: CpuThreadCountControl::PerCallUpperBound,
placement: CpuPlacementControl::EngineWorkers,
worker_local_sequential: true,
accepts_sequential: true,
accepts_outer: true,
accepts_inner: true,
}
}
#[cfg(test)]
fn controlled_external_capabilities(
thread_count: CpuThreadCountControl,
) -> CpuProviderExecutionCapabilities {
CpuProviderExecutionCapabilities {
thread_count,
placement: CpuPlacementControl::ExternalWorkers,
worker_local_sequential: true,
accepts_sequential: true,
accepts_outer: true,
accepts_inner: true,
}
}
#[cfg(any(test, feature = "cpu-blas"))]
fn uncontrolled_external_capabilities() -> CpuProviderExecutionCapabilities {
CpuProviderExecutionCapabilities {
thread_count: CpuThreadCountControl::GlobalOrUncontrolled,
placement: CpuPlacementControl::ExternalWorkers,
worker_local_sequential: false,
accepts_sequential: false,
accepts_outer: false,
accepts_inner: true,
}
}
#[cfg(any(test, not(feature = "cpu-blas")))]
pub(crate) fn serial_capabilities() -> CpuProviderExecutionCapabilities {
CpuProviderExecutionCapabilities {
thread_count: CpuThreadCountControl::Sequential,
placement: CpuPlacementControl::CallingThread,
worker_local_sequential: true,
accepts_sequential: true,
accepts_outer: true,
accepts_inner: true,
}
}
#[cfg(any(test, feature = "cpu-blas"))]
pub(crate) fn builtin_blas_execution_capabilities() -> CpuProviderExecutionCapabilities {
uncontrolled_external_capabilities()
}
pub(crate) fn validate_provider_for_caller_managed_domain(
capabilities: CpuProviderExecutionCapabilities,
thread_budget: NonZeroUsize,
) -> Result<(), CpuProviderDomainError> {
if enforced_provider_thread_limit(capabilities.thread_count, thread_budget).is_none() {
return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
thread_budget: thread_budget.get(),
control: capabilities.thread_count,
});
}
match capabilities.placement {
CpuPlacementControl::EngineWorkers | CpuPlacementControl::CallingThread => Ok(()),
CpuPlacementControl::ExternalWorkers | CpuPlacementControl::None => Err(
CpuProviderDomainError::CallerManagedPlacementNotEnforceable {
thread_budget: thread_budget.get(),
placement: capabilities.placement,
},
),
}
}
pub(crate) fn validate_provider_for_domain(
capabilities: CpuProviderExecutionCapabilities,
thread_budget: NonZeroUsize,
placement_guarantee: CpuPlacementGuarantee,
domain_cpus: &CpuSet,
process_allowed_cpus: &CpuSet,
) -> Result<(), CpuProviderDomainError> {
if enforced_provider_thread_limit(capabilities.thread_count, thread_budget).is_none() {
return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
thread_budget: thread_budget.get(),
control: capabilities.thread_count,
});
}
match capabilities.placement {
CpuPlacementControl::EngineWorkers | CpuPlacementControl::CallingThread => Ok(()),
CpuPlacementControl::ExternalWorkers => {
if thread_budget.get() == 1 && capabilities.worker_local_sequential {
return Ok(());
}
if placement_guarantee == CpuPlacementGuarantee::AdvisoryDeclared
|| domain_cpus == process_allowed_cpus
{
return Ok(());
}
Err(CpuProviderDomainError::PlacementNotEnforceable {
thread_budget: thread_budget.get(),
placement: capabilities.placement,
guarantee: placement_guarantee,
})
}
CpuPlacementControl::None => {
if placement_guarantee == CpuPlacementGuarantee::AdvisoryDeclared {
Ok(())
} else {
Err(CpuProviderDomainError::PlacementNotEnforceable {
thread_budget: thread_budget.get(),
placement: capabilities.placement,
guarantee: placement_guarantee,
})
}
}
}
}
fn enforced_provider_thread_limit(
control: CpuThreadCountControl,
thread_budget: NonZeroUsize,
) -> Option<NonZeroUsize> {
match control {
CpuThreadCountControl::Sequential | CpuThreadCountControl::BinaryClampToOne => {
NonZeroUsize::new(1)
}
CpuThreadCountControl::PerCallUpperBound => Some(thread_budget),
CpuThreadCountControl::GlobalOrUncontrolled => None,
}
}
#[cfg(test)]
mod tests;