1use std::num::NonZeroUsize;
2
3use thiserror::Error;
4
5use crate::ParallelMode;
6
7#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
19pub enum CpuThreadCountControl {
20 Sequential,
22 PerCallUpperBound,
24 BinaryClampToOne,
30 #[default]
32 GlobalOrUncontrolled,
33}
34
35#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
47pub enum CpuPlacementControl {
48 EngineWorkers,
50 CallingThread,
52 ExternalWorkers,
54 #[default]
56 None,
57}
58
59#[derive(Clone, Copy, Debug, Eq, PartialEq)]
81pub struct CpuProviderExecutionCapabilities {
82 pub thread_count: CpuThreadCountControl,
84 pub placement: CpuPlacementControl,
86 pub worker_local_sequential: bool,
88 pub accepts_sequential: bool,
90 pub accepts_outer: bool,
92 pub accepts_inner: bool,
94}
95
96impl Default for CpuProviderExecutionCapabilities {
97 fn default() -> Self {
98 Self {
99 thread_count: CpuThreadCountControl::GlobalOrUncontrolled,
100 placement: CpuPlacementControl::None,
101 worker_local_sequential: false,
102 accepts_sequential: false,
103 accepts_outer: false,
104 accepts_inner: true,
105 }
106 }
107}
108
109#[derive(Clone, Copy, Debug, Eq, Error, PartialEq)]
122pub enum CpuProviderDomainError {
123 #[error(
125 "provider thread-count control {control:?} cannot enforce thread budget {thread_budget}"
126 )]
127 ThreadCountNotEnforceable {
128 thread_budget: usize,
130 control: CpuThreadCountControl,
132 },
133 #[error(
135 "provider placement control {placement:?} can leave the caller-managed executor for thread budget {thread_budget}"
136 )]
137 CallerManagedPlacementNotEnforceable {
138 thread_budget: usize,
140 placement: CpuPlacementControl,
142 },
143 #[error("provider cannot honor requested CPU parallel mode {mode:?}")]
145 ParallelModeNotSupported {
146 mode: ParallelMode,
148 },
149}
150
151impl CpuProviderExecutionCapabilities {
152 pub(crate) fn accepts_mode(self, mode: ParallelMode) -> bool {
153 match mode {
154 ParallelMode::Sequential => self.accepts_sequential && self.worker_local_sequential,
155 ParallelMode::Outer => self.accepts_outer && self.worker_local_sequential,
156 ParallelMode::Inner => self.accepts_inner,
157 }
158 }
159}
160
161#[cfg(test)]
162#[derive(Clone, Copy, Debug, Eq, PartialEq)]
163pub(crate) enum OpenBlasParallelism {
164 Sequential,
165 Pthread,
166 OpenMp,
167 Unknown,
168}
169
170#[cfg(test)]
171#[derive(Clone, Copy, Debug, Eq, PartialEq)]
172pub(crate) struct OpenBlasProbe {
173 pub(crate) parallelism: OpenBlasParallelism,
174 pub(crate) process_global_set_restore_wired: bool,
175}
176
177#[cfg(test)]
178#[derive(Clone, Copy, Debug, Eq, PartialEq)]
179pub(crate) struct AccelerateProbe {
180 pub(crate) binary_thread_local_control_wired: bool,
181}
182
183#[cfg(test)]
190#[derive(Clone, Copy, Debug, Eq, PartialEq)]
191pub(crate) enum CpuProviderProbe {
192 FaerOrNative,
193 Mkl { thread_local_setter_wired: bool },
194 OpenBlas(OpenBlasProbe),
195 Accelerate(AccelerateProbe),
196 ArmPlOpenMp,
197 ArmPlSerial,
198 NvplSerial,
199 UnknownBlas,
200 Injected(Option<CpuProviderExecutionCapabilities>),
201}
202
203#[cfg(test)]
204pub(crate) fn classify_provider(probe: CpuProviderProbe) -> CpuProviderExecutionCapabilities {
205 match probe {
206 CpuProviderProbe::FaerOrNative => engine_worker_capabilities(),
207 CpuProviderProbe::Mkl {
208 thread_local_setter_wired: true,
209 } => controlled_external_capabilities(CpuThreadCountControl::PerCallUpperBound),
210 CpuProviderProbe::Mkl {
211 thread_local_setter_wired: false,
212 }
213 | CpuProviderProbe::ArmPlOpenMp => uncontrolled_external_capabilities(),
214 CpuProviderProbe::OpenBlas(probe) => classify_openblas(probe),
215 CpuProviderProbe::Accelerate(probe) => classify_accelerate(probe),
216 CpuProviderProbe::ArmPlSerial | CpuProviderProbe::NvplSerial => serial_capabilities(),
217 CpuProviderProbe::UnknownBlas | CpuProviderProbe::Injected(None) => {
218 CpuProviderExecutionCapabilities::default()
219 }
220 CpuProviderProbe::Injected(Some(capabilities)) => capabilities,
221 }
222}
223
224#[cfg(test)]
225fn classify_openblas(probe: OpenBlasProbe) -> CpuProviderExecutionCapabilities {
226 match (probe.parallelism, probe.process_global_set_restore_wired) {
227 (OpenBlasParallelism::Sequential, _) => serial_capabilities(),
228 (OpenBlasParallelism::Pthread | OpenBlasParallelism::OpenMp, _) => {
229 uncontrolled_external_capabilities()
230 }
231 (OpenBlasParallelism::Unknown, _) => CpuProviderExecutionCapabilities::default(),
232 }
233}
234
235#[cfg(test)]
236fn classify_accelerate(probe: AccelerateProbe) -> CpuProviderExecutionCapabilities {
237 if probe.binary_thread_local_control_wired {
238 controlled_external_capabilities(CpuThreadCountControl::BinaryClampToOne)
239 } else {
240 uncontrolled_external_capabilities()
241 }
242}
243
244pub(crate) fn engine_worker_capabilities() -> CpuProviderExecutionCapabilities {
245 CpuProviderExecutionCapabilities {
246 thread_count: CpuThreadCountControl::PerCallUpperBound,
247 placement: CpuPlacementControl::EngineWorkers,
248 worker_local_sequential: true,
249 accepts_sequential: true,
250 accepts_outer: true,
251 accepts_inner: true,
252 }
253}
254
255#[cfg(test)]
256fn controlled_external_capabilities(
257 thread_count: CpuThreadCountControl,
258) -> CpuProviderExecutionCapabilities {
259 CpuProviderExecutionCapabilities {
260 thread_count,
261 placement: CpuPlacementControl::ExternalWorkers,
262 worker_local_sequential: true,
263 accepts_sequential: true,
264 accepts_outer: true,
265 accepts_inner: true,
266 }
267}
268
269#[cfg(any(test, feature = "cpu-blas"))]
270pub(crate) fn uncontrolled_external_capabilities() -> CpuProviderExecutionCapabilities {
271 CpuProviderExecutionCapabilities {
272 thread_count: CpuThreadCountControl::GlobalOrUncontrolled,
273 placement: CpuPlacementControl::ExternalWorkers,
274 worker_local_sequential: false,
275 accepts_sequential: false,
276 accepts_outer: false,
277 accepts_inner: true,
278 }
279}
280
281#[cfg(any(test, not(feature = "cpu-blas")))]
282pub(crate) fn serial_capabilities() -> CpuProviderExecutionCapabilities {
283 CpuProviderExecutionCapabilities {
284 thread_count: CpuThreadCountControl::Sequential,
285 placement: CpuPlacementControl::CallingThread,
286 worker_local_sequential: true,
287 accepts_sequential: true,
288 accepts_outer: true,
289 accepts_inner: true,
290 }
291}
292
293#[cfg(any(test, feature = "cpu-blas"))]
298pub(crate) fn builtin_blas_execution_capabilities() -> CpuProviderExecutionCapabilities {
299 uncontrolled_external_capabilities()
300}
301
302pub(crate) fn validate_provider_for_caller_managed_domain(
303 capabilities: CpuProviderExecutionCapabilities,
304 thread_budget: NonZeroUsize,
305) -> Result<(), CpuProviderDomainError> {
306 if enforced_provider_thread_limit(capabilities.thread_count, thread_budget).is_none() {
307 return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
308 thread_budget: thread_budget.get(),
309 control: capabilities.thread_count,
310 });
311 }
312 match capabilities.placement {
313 CpuPlacementControl::EngineWorkers | CpuPlacementControl::CallingThread => Ok(()),
314 CpuPlacementControl::ExternalWorkers | CpuPlacementControl::None => Err(
315 CpuProviderDomainError::CallerManagedPlacementNotEnforceable {
316 thread_budget: thread_budget.get(),
317 placement: capabilities.placement,
318 },
319 ),
320 }
321}
322
323pub(crate) fn validate_provider_for_domain(
331 capabilities: CpuProviderExecutionCapabilities,
332 thread_budget: NonZeroUsize,
333) -> Result<(), CpuProviderDomainError> {
334 if enforced_provider_thread_limit(capabilities.thread_count, thread_budget).is_none() {
335 return Err(CpuProviderDomainError::ThreadCountNotEnforceable {
336 thread_budget: thread_budget.get(),
337 control: capabilities.thread_count,
338 });
339 }
340 Ok(())
341}
342
343fn enforced_provider_thread_limit(
344 control: CpuThreadCountControl,
345 thread_budget: NonZeroUsize,
346) -> Option<NonZeroUsize> {
347 match control {
348 CpuThreadCountControl::Sequential | CpuThreadCountControl::BinaryClampToOne => {
349 NonZeroUsize::new(1)
350 }
351 CpuThreadCountControl::PerCallUpperBound => Some(thread_budget),
352 CpuThreadCountControl::GlobalOrUncontrolled => None,
353 }
354}
355
356#[cfg(test)]
357mod tests;