Skip to main content

strided_basic/
execution_policy.rs

1use core::cell::Cell;
2use core::num::NonZeroUsize;
3
4/// Controls how a strided operation may use CPU threads.
5///
6/// This policy controls fanout only. It does not create a Rayon pool or choose
7/// CPU placement; callers that need placement control should install the
8/// operation in their chosen executor and then select [`Self::Sequential`] or
9/// [`Self::Rayon`].
10///
11/// The policy applies to fanout owned by `strided-kernel`. A bounded worker
12/// partition runs nested strided operations sequentially so nested operations
13/// cannot multiply the outer budget. The scope is worker-local, not a
14/// task-local Rayon context, and is not propagated to threads or Rayon tasks
15/// created by user callbacks.
16///
17/// User callbacks that rely on policy isolation must not enter their own Rayon
18/// scheduling or yield boundary (`join`, `scope`, `spawn`, and similar APIs).
19/// At such a callback-owned boundary, Rayon may execute unrelated work on the
20/// waiting worker while its worker-local policy is active. `strided-kernel`
21/// suspends the policy at every scheduler boundary it owns, but cannot do so
22/// for arbitrary scheduling performed inside a callback.
23#[derive(Clone, Copy, Debug, Eq, PartialEq)]
24pub enum ExecutionPolicy {
25    /// Preserve the compatibility behavior of using all threads in the current
26    /// installed Rayon pool, or the global pool when no pool is installed.
27    ///
28    /// Explicit runtimes should prefer [`Self::Sequential`] or [`Self::Rayon`].
29    AmbientRayon,
30    /// Execute entirely on the calling thread with zero Rayon fanout.
31    Sequential,
32    /// Use the currently installed Rayon pool while limiting the operation to
33    /// at most `max_threads` concurrent partitions.
34    ///
35    /// Without the `parallel` feature, operations remain sequential.
36    Rayon { max_threads: NonZeroUsize },
37}
38
39thread_local! {
40    static ACTIVE_POLICY: Cell<ExecutionPolicy> = const { Cell::new(ExecutionPolicy::AmbientRayon) };
41    static ACTIVE_FANOUT: Cell<bool> = const { Cell::new(false) };
42}
43
44fn restrict(outer: ExecutionPolicy, inner: ExecutionPolicy) -> ExecutionPolicy {
45    match (outer, inner) {
46        (ExecutionPolicy::Sequential, _) | (_, ExecutionPolicy::Sequential) => {
47            ExecutionPolicy::Sequential
48        }
49        (ExecutionPolicy::AmbientRayon, policy) | (policy, ExecutionPolicy::AmbientRayon) => policy,
50        (
51            ExecutionPolicy::Rayon { max_threads: outer },
52            ExecutionPolicy::Rayon { max_threads: inner },
53        ) => ExecutionPolicy::Rayon {
54            max_threads: outer.min(inner),
55        },
56    }
57}
58
59#[derive(Clone, Copy)]
60struct ExecutionState {
61    policy: ExecutionPolicy,
62    fanout_active: bool,
63}
64
65struct StateGuard {
66    previous: ExecutionState,
67}
68
69impl Drop for StateGuard {
70    fn drop(&mut self) {
71        set_state(self.previous);
72    }
73}
74
75fn state() -> ExecutionState {
76    ExecutionState {
77        policy: ACTIVE_POLICY.with(Cell::get),
78        fanout_active: ACTIVE_FANOUT.with(Cell::get),
79    }
80}
81
82fn set_state(state: ExecutionState) {
83    ACTIVE_POLICY.with(|active| active.set(state.policy));
84    ACTIVE_FANOUT.with(|active| active.set(state.fanout_active));
85}
86
87#[cfg(feature = "parallel")]
88fn with_state<R>(next: ExecutionState, operation: impl FnOnce() -> R) -> R {
89    let previous = state();
90    set_state(next);
91    let _guard = StateGuard { previous };
92    operation()
93}
94
95/// Execute `operation` under an explicit strided-kernel execution policy.
96///
97/// Nested scopes combine conservatively: sequential execution dominates and
98/// nested Rayon budgets use the smaller limit. The previous policy is restored
99/// even if `operation` panics.
100///
101/// The contract covers strided operations and their library-owned fanout. It
102/// does not turn Rayon worker-local state into task-local state. In particular,
103/// callbacks should not invoke Rayon scheduling or yielding APIs if they depend
104/// on isolation from unrelated work; see [`ExecutionPolicy`] for details.
105///
106/// # Examples
107///
108/// ```rust
109/// use strided_basic::{
110///     map_into, with_execution_policy, ExecutionPolicy, StridedArray,
111/// };
112///
113/// let source = StridedArray::<f64>::from_fn_col_major(&[3], |index| index[0] as f64);
114/// let mut destination = StridedArray::<f64>::col_major(&[3]);
115/// with_execution_policy(ExecutionPolicy::Sequential, || {
116///     map_into(&mut destination.view_mut(), &source.view(), |value| value + 1.0)
117///         .unwrap();
118/// });
119/// assert_eq!(destination.into_data(), vec![1.0, 2.0, 3.0]);
120/// ```
121#[inline]
122pub fn with_execution_policy<R>(policy: ExecutionPolicy, operation: impl FnOnce() -> R) -> R {
123    let policy = match policy {
124        ExecutionPolicy::AmbientRayon => return operation(),
125        policy => policy,
126    };
127    let previous = state();
128    set_state(ExecutionState {
129        policy: restrict(previous.policy, policy),
130        fanout_active: previous.fanout_active,
131    });
132    let _guard = StateGuard { previous };
133    operation()
134}
135
136#[cfg(feature = "parallel")]
137pub(crate) fn active_policy() -> ExecutionPolicy {
138    ACTIVE_POLICY.with(Cell::get)
139}
140
141#[cfg(feature = "parallel")]
142pub(crate) fn fanout_active() -> bool {
143    ACTIVE_FANOUT.with(Cell::get)
144}
145
146#[cfg(feature = "parallel")]
147#[inline(always)]
148pub(crate) fn with_owned_execution<R>(
149    policy: ExecutionPolicy,
150    fanout_active: bool,
151    operation: impl FnOnce() -> R,
152) -> R {
153    match policy {
154        ExecutionPolicy::AmbientRayon => operation(),
155        ExecutionPolicy::Sequential | ExecutionPolicy::Rayon { .. } => {
156            let previous = state();
157            with_state(
158                ExecutionState {
159                    policy: restrict(previous.policy, policy),
160                    fanout_active: previous.fanout_active || fanout_active,
161                },
162                operation,
163            )
164        }
165    }
166}
167
168#[cfg(feature = "parallel")]
169pub(crate) fn with_scheduler_suspended<R>(operation: impl FnOnce() -> R) -> R {
170    with_state(
171        ExecutionState {
172            policy: ExecutionPolicy::AmbientRayon,
173            fanout_active: false,
174        },
175        operation,
176    )
177}
178
179#[cfg(feature = "parallel")]
180pub(crate) fn permutation_copy_parallel_eligible(
181    policy: ExecutionPolicy,
182    fanout_active: bool,
183    current_pool_threads: usize,
184) -> bool {
185    if fanout_active || current_pool_threads <= 1 {
186        return false;
187    }
188    match policy {
189        ExecutionPolicy::AmbientRayon => true,
190        ExecutionPolicy::Sequential => false,
191        ExecutionPolicy::Rayon { max_threads } => current_pool_threads <= max_threads.get(),
192    }
193}
194
195#[cfg(feature = "parallel")]
196pub fn rayon_threads() -> usize {
197    if fanout_active() {
198        return 1;
199    }
200    match active_policy() {
201        ExecutionPolicy::Sequential => 1,
202        ExecutionPolicy::AmbientRayon => crate::threading::current_pool_threads(),
203        ExecutionPolicy::Rayon { max_threads } => {
204            crate::threading::current_pool_threads().min(max_threads.get())
205        }
206    }
207}
208
209#[cfg(test)]
210#[path = "execution_policy/tests/default_tests.rs"]
211mod default_tests;
212
213#[cfg(all(test, feature = "parallel"))]
214#[path = "execution_policy/tests/tests.rs"]
215mod tests;