Skip to main content

tenferro_cpu/
batch_policy.rs

1//! Overridable batch execution strategy for batched CPU operations (#1938 D9).
2//!
3//! A batched operation (strided-batched or grouped GEMM, packed LU/solve) runs
4//! its independent items one of several ways. [`CpuBatchStrategy::Auto`] keeps
5//! the backend heuristics; every other strategy forces one route and fails with
6//! a typed error when that route is not available, instead of falling back.
7//!
8//! Precedence is per-operation > scoped override > backend default. The backend
9//! default is set with [`crate::CpuBackend::with_batch_policy`]; a scoped
10//! override, which also expresses a per-operation choice when it wraps a single
11//! call, is [`crate::with_batch_policy`]. Scopes nest and the
12//! innermost wins; each restores the previous policy on return, error and
13//! unwind. Nothing here mutates process-global state.
14//!
15//! # Examples
16//!
17//! ```rust
18//! use tenferro_cpu::{CpuBackend, CpuBatchPolicy, CpuBatchStrategy};
19//!
20//! let backend = CpuBackend::with_threads(1)?
21//!     .with_batch_policy(CpuBatchPolicy::new(CpuBatchStrategy::Sequential));
22//! assert_eq!(backend.batch_policy().strategy(), CpuBatchStrategy::Sequential);
23//! # Ok::<(), Box<dyn std::error::Error>>(())
24//! ```
25
26/// How the items of one batched operation are executed.
27///
28/// # Examples
29///
30/// ```rust
31/// use tenferro_cpu::CpuBatchStrategy;
32///
33/// assert_eq!(CpuBatchStrategy::default(), CpuBatchStrategy::Auto);
34/// ```
35#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
36#[non_exhaustive]
37pub enum CpuBatchStrategy {
38    /// Choose a supported route from the thresholds and the provider.
39    #[default]
40    Auto,
41    /// Sequential batch, sequential items: no parallelism at all.
42    Sequential,
43    /// Outer-parallel batch: tenferro lanes, each running items sequentially.
44    OuterParallel,
45    /// Sequential batch whose items may use the provider's own parallelism.
46    ProviderItems,
47    /// One vendor batch call for the whole batch, with no tenferro fan-out.
48    WholeBatchVendor,
49}
50
51/// Thresholds that [`CpuBatchStrategy::Auto`] applies.
52///
53/// The defaults are the values the backend used before these knobs existed;
54/// they are not tuned results.
55///
56/// # Examples
57///
58/// ```rust
59/// use tenferro_cpu::CpuBatchThresholds;
60///
61/// let thresholds = CpuBatchThresholds::default()
62///     .with_vendor_batch_max_item_dim(8)
63///     .with_outer_min_items(4);
64/// assert_eq!(thresholds.vendor_batch_max_item_dim(), 8);
65/// assert_eq!(thresholds.outer_min_items(), 4);
66/// assert_eq!(thresholds.outer_min_items_per_lane(), 1);
67/// ```
68#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
69pub struct CpuBatchThresholds {
70    vendor_batch_max_item_dim: usize,
71    outer_min_items: usize,
72    outer_min_items_per_lane: usize,
73    lane_item_overhead_ns: usize,
74    lane_muladds_per_ns: usize,
75    lane_min_work_ns: usize,
76}
77
78impl Default for CpuBatchThresholds {
79    fn default() -> Self {
80        Self {
81            // Per-item work: the BLAS provider's measured small-job cutoff.
82            vendor_batch_max_item_dim: 16,
83            // Total batch work: fan-out needs more than one item.
84            outer_min_items: 2,
85            // Chunk granularity: every lane must receive at least one item.
86            outer_min_items_per_lane: 1,
87            // Lane cost model, fitted on an AMD EPYC host (1..16 threads,
88            // faer, f64): one GEMM item costs about 50 ns of per-call work plus
89            // 1 ns per 16 multiply-adds, and a lane pays off only with about
90            // 8 us of estimated work. A pure multiply-add cutoff missed large
91            // batches of tiny items; too many short lanes made 16 threads
92            // slower than one.
93            lane_item_overhead_ns: 50,
94            lane_muladds_per_ns: 16,
95            lane_min_work_ns: 8_000,
96        }
97    }
98}
99
100impl CpuBatchThresholds {
101    /// Per-item work: the largest `m`, `n` and `k` for which `Auto` lets a
102    /// vendor batch call (`cblas_?gemm_batch`) handle a grouped GEMM batch.
103    ///
104    /// Strided batched contractions do not use this cutoff: under `Auto` they
105    /// keep one provider GEMM per item, and reach the vendor batch call only
106    /// through [`CpuBatchStrategy::WholeBatchVendor`].
107    ///
108    /// # Examples
109    ///
110    /// ```rust
111    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().vendor_batch_max_item_dim(), 16);
112    /// ```
113    #[must_use]
114    pub fn vendor_batch_max_item_dim(&self) -> usize {
115        self.vendor_batch_max_item_dim
116    }
117
118    /// Total batch work: the fewest items for which `Auto` fans out.
119    ///
120    /// # Examples
121    ///
122    /// ```rust
123    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().outer_min_items(), 2);
124    /// ```
125    #[must_use]
126    pub fn outer_min_items(&self) -> usize {
127        self.outer_min_items
128    }
129
130    /// Chunk granularity: the fewest items each lane must receive before
131    /// `Auto` fans a batch out over every lane.
132    ///
133    /// # Examples
134    ///
135    /// ```rust
136    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().outer_min_items_per_lane(), 1);
137    /// ```
138    #[must_use]
139    pub fn outer_min_items_per_lane(&self) -> usize {
140        self.outer_min_items_per_lane
141    }
142
143    /// Lane cost model: the estimated fixed cost of one GEMM item, in
144    /// nanoseconds, that a lane must amortize.
145    ///
146    /// # Examples
147    ///
148    /// ```rust
149    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().lane_item_overhead_ns(), 50);
150    /// ```
151    #[must_use]
152    pub fn lane_item_overhead_ns(&self) -> usize {
153        self.lane_item_overhead_ns
154    }
155
156    /// Lane cost model: estimated multiply-adds per nanosecond of one lane.
157    ///
158    /// # Examples
159    ///
160    /// ```rust
161    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().lane_muladds_per_ns(), 16);
162    /// ```
163    #[must_use]
164    pub fn lane_muladds_per_ns(&self) -> usize {
165        self.lane_muladds_per_ns
166    }
167
168    /// Lane cost model: the least estimated work, in nanoseconds, each lane
169    /// must receive before `Auto` splits a batch inside an entered session.
170    /// Zero lets the item thresholds alone decide.
171    ///
172    /// # Examples
173    ///
174    /// ```rust
175    /// assert_eq!(tenferro_cpu::CpuBatchThresholds::default().lane_min_work_ns(), 8_000);
176    /// ```
177    #[must_use]
178    pub fn lane_min_work_ns(&self) -> usize {
179        self.lane_min_work_ns
180    }
181
182    /// Return these thresholds with a different per-item lane overhead.
183    ///
184    /// # Examples
185    ///
186    /// ```rust
187    /// use tenferro_cpu::CpuBatchThresholds;
188    /// let thresholds = CpuBatchThresholds::default().with_lane_item_overhead_ns(100);
189    /// assert_eq!(thresholds.lane_item_overhead_ns(), 100);
190    /// ```
191    #[must_use]
192    pub fn with_lane_item_overhead_ns(mut self, nanoseconds: usize) -> Self {
193        self.lane_item_overhead_ns = nanoseconds;
194        self
195    }
196
197    /// Return these thresholds with a different lane throughput estimate.
198    /// Zero is treated as one multiply-add per nanosecond.
199    ///
200    /// # Examples
201    ///
202    /// ```rust
203    /// use tenferro_cpu::CpuBatchThresholds;
204    /// let thresholds = CpuBatchThresholds::default().with_lane_muladds_per_ns(32);
205    /// assert_eq!(thresholds.lane_muladds_per_ns(), 32);
206    /// ```
207    #[must_use]
208    pub fn with_lane_muladds_per_ns(mut self, muladds: usize) -> Self {
209        self.lane_muladds_per_ns = muladds;
210        self
211    }
212
213    /// Return these thresholds with a different minimum work per lane.
214    ///
215    /// # Examples
216    ///
217    /// ```rust
218    /// use tenferro_cpu::CpuBatchThresholds;
219    /// let thresholds = CpuBatchThresholds::default().with_lane_min_work_ns(0);
220    /// assert_eq!(thresholds.lane_min_work_ns(), 0);
221    /// ```
222    #[must_use]
223    pub fn with_lane_min_work_ns(mut self, nanoseconds: usize) -> Self {
224        self.lane_min_work_ns = nanoseconds;
225        self
226    }
227
228    /// Estimated cost of one `m x n x k` GEMM item under the lane cost model.
229    pub(crate) fn lane_item_ns(&self, m: usize, n: usize, k: usize) -> usize {
230        let muladds = m.saturating_mul(n).saturating_mul(k);
231        self.lane_item_overhead_ns
232            .saturating_add(muladds / self.lane_muladds_per_ns.max(1))
233    }
234
235    /// The lanes `Auto` uses for `items` jobs of `total_ns` estimated work on
236    /// `threads` threads, or `None` when fewer than two lanes would each
237    /// receive [`Self::lane_min_work_ns`] and pass [`Self::fans_out`].
238    pub(crate) fn auto_lanes(
239        &self,
240        items: usize,
241        total_ns: usize,
242        threads: usize,
243    ) -> Option<usize> {
244        let by_work = match self.lane_min_work_ns {
245            0 => usize::MAX,
246            min => total_ns / min,
247        };
248        let lanes = threads.min(items).min(by_work);
249        (lanes >= 2 && self.fans_out(items, lanes)).then_some(lanes)
250    }
251
252    /// Return these thresholds with a different vendor-batch item limit.
253    ///
254    /// # Examples
255    ///
256    /// ```rust
257    /// use tenferro_cpu::CpuBatchThresholds;
258    /// let thresholds = CpuBatchThresholds::default().with_vendor_batch_max_item_dim(0);
259    /// assert_eq!(thresholds.vendor_batch_max_item_dim(), 0);
260    /// ```
261    #[must_use]
262    pub fn with_vendor_batch_max_item_dim(mut self, limit: usize) -> Self {
263        self.vendor_batch_max_item_dim = limit;
264        self
265    }
266
267    /// Return these thresholds with a different minimum fan-out batch size.
268    ///
269    /// # Examples
270    ///
271    /// ```rust
272    /// use tenferro_cpu::CpuBatchThresholds;
273    /// let thresholds = CpuBatchThresholds::default().with_outer_min_items(64);
274    /// assert_eq!(thresholds.outer_min_items(), 64);
275    /// ```
276    #[must_use]
277    pub fn with_outer_min_items(mut self, items: usize) -> Self {
278        self.outer_min_items = items;
279        self
280    }
281
282    /// Return these thresholds with a different minimum chunk per lane.
283    ///
284    /// # Examples
285    ///
286    /// ```rust
287    /// use tenferro_cpu::CpuBatchThresholds;
288    /// let thresholds = CpuBatchThresholds::default().with_outer_min_items_per_lane(8);
289    /// assert_eq!(thresholds.outer_min_items_per_lane(), 8);
290    /// ```
291    #[must_use]
292    pub fn with_outer_min_items_per_lane(mut self, items: usize) -> Self {
293        self.outer_min_items_per_lane = items;
294        self
295    }
296
297    /// Whether `Auto` fans `items` out over `lanes` lanes: more than one lane,
298    /// at least [`Self::outer_min_items`] items, and at least
299    /// [`Self::outer_min_items_per_lane`] items for every lane.
300    ///
301    /// # Examples
302    ///
303    /// ```rust
304    /// use tenferro_cpu::CpuBatchThresholds;
305    ///
306    /// let thresholds = CpuBatchThresholds::default();
307    /// assert!(thresholds.fans_out(8, 4));
308    /// assert!(!thresholds.fans_out(3, 4));
309    /// assert!(!thresholds.fans_out(8, 1));
310    /// ```
311    #[must_use]
312    pub fn fans_out(&self, items: usize, lanes: usize) -> bool {
313        lanes > 1
314            && items >= self.outer_min_items
315            && items >= lanes.saturating_mul(self.outer_min_items_per_lane)
316    }
317
318    /// Whether `Auto` may hand every GEMM of dimensions `dims` to one vendor
319    /// batch call.
320    pub(crate) fn auto_uses_vendor_batch(
321        &self,
322        dims: impl IntoIterator<Item = [usize; 3]>,
323    ) -> bool {
324        let limit = self.vendor_batch_max_item_dim;
325        let mut count = 0usize;
326        for [m, n, k] in dims {
327            if m > limit || n > limit || k > limit {
328                return false;
329            }
330            count += 1;
331        }
332        count > 1
333    }
334}
335
336/// A batch strategy together with the thresholds `Auto` uses.
337///
338/// # Examples
339///
340/// ```rust
341/// use tenferro_cpu::{CpuBatchPolicy, CpuBatchStrategy, CpuBatchThresholds};
342///
343/// let policy = CpuBatchPolicy::default()
344///     .with_thresholds(CpuBatchThresholds::default().with_outer_min_items(8));
345/// assert_eq!(policy.strategy(), CpuBatchStrategy::Auto);
346/// assert_eq!(policy.thresholds().outer_min_items(), 8);
347/// ```
348#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
349pub struct CpuBatchPolicy {
350    strategy: CpuBatchStrategy,
351    thresholds: CpuBatchThresholds,
352}
353
354impl CpuBatchPolicy {
355    /// A policy with `strategy` and the default thresholds.
356    ///
357    /// # Examples
358    ///
359    /// ```rust
360    /// use tenferro_cpu::{CpuBatchPolicy, CpuBatchStrategy};
361    /// let policy = CpuBatchPolicy::new(CpuBatchStrategy::WholeBatchVendor);
362    /// assert_eq!(policy.strategy(), CpuBatchStrategy::WholeBatchVendor);
363    /// ```
364    #[must_use]
365    pub fn new(strategy: CpuBatchStrategy) -> Self {
366        Self {
367            strategy,
368            thresholds: CpuBatchThresholds::default(),
369        }
370    }
371
372    /// Return this policy with different thresholds.
373    ///
374    /// # Examples
375    ///
376    /// ```rust
377    /// use tenferro_cpu::{CpuBatchPolicy, CpuBatchThresholds};
378    /// let policy = CpuBatchPolicy::default().with_thresholds(CpuBatchThresholds::default());
379    /// assert_eq!(policy, CpuBatchPolicy::default());
380    /// ```
381    #[must_use]
382    pub fn with_thresholds(mut self, thresholds: CpuBatchThresholds) -> Self {
383        self.thresholds = thresholds;
384        self
385    }
386
387    /// The selected strategy.
388    ///
389    /// # Examples
390    ///
391    /// ```rust
392    /// use tenferro_cpu::{CpuBatchPolicy, CpuBatchStrategy};
393    /// assert_eq!(CpuBatchPolicy::default().strategy(), CpuBatchStrategy::Auto);
394    /// ```
395    #[must_use]
396    pub fn strategy(&self) -> CpuBatchStrategy {
397        self.strategy
398    }
399
400    /// The thresholds `Auto` applies.
401    ///
402    /// # Examples
403    ///
404    /// ```rust
405    /// use tenferro_cpu::{CpuBatchPolicy, CpuBatchThresholds};
406    /// assert_eq!(CpuBatchPolicy::default().thresholds(), CpuBatchThresholds::default());
407    /// ```
408    #[must_use]
409    pub fn thresholds(&self) -> CpuBatchThresholds {
410        self.thresholds
411    }
412}
413
414/// Run `f` on `session` with `policy` as the effective batch policy.
415///
416/// This is the scoped override of the precedence per-operation > scoped >
417/// backend default; wrapping a single call expresses a per-operation choice.
418/// The previous policy is restored when `f` returns, returns an error or
419/// unwinds, and scopes nest with the innermost winning.
420///
421/// # Examples
422///
423/// ```rust
424/// use tenferro_cpu::{with_batch_policy, CpuBackend, CpuBatchPolicy, CpuBatchStrategy};
425/// use tenferro_tensor::{BackendSessionHost, DotGeneralConfig, Tensor, TensorRead};
426///
427/// let mut backend = CpuBackend::with_threads(1)?;
428/// let lhs = Tensor::from_vec_col_major(vec![2, 2, 3], vec![1.0_f64; 12])?;
429/// let rhs = Tensor::from_vec_col_major(vec![2, 2, 3], vec![2.0_f64; 12])?;
430/// let config = DotGeneralConfig {
431///     lhs_contracting_dims: [1].as_slice().into(),
432///     rhs_contracting_dims: [0].as_slice().into(),
433///     lhs_batch_dims: [2].as_slice().into(),
434///     rhs_batch_dims: [2].as_slice().into(),
435/// };
436/// let product = backend.with_backend_session(|session| {
437///     // ProviderItems runs one provider GEMM per item with every provider;
438///     // a forced Sequential is rejected by providers with their own threading.
439///     with_batch_policy(session, CpuBatchPolicy::new(CpuBatchStrategy::ProviderItems), |session| {
440///         session.dot_general_read(
441///             TensorRead::from_tensor(&lhs),
442///             TensorRead::from_tensor(&rhs),
443///             &config,
444///         )
445///     })
446/// })???;
447/// assert_eq!(product.as_slice::<f64>()?, &[4.0; 12]);
448/// # Ok::<(), Box<dyn std::error::Error>>(())
449/// ```
450///
451/// # Errors
452///
453/// Returns [`tenferro_tensor::Error::Unsupported`] without running `f` when
454/// `session` is neither a CPU execution session nor a session that forwards
455/// one through [`tenferro_tensor::BackendSession::native_session`] (a CUDA or
456/// other custom session). `f`'s own result is returned unchanged inside `Ok`.
457///
458/// # Panics
459///
460/// A panic in `f` propagates after the previous policy is restored.
461pub fn with_batch_policy<R>(
462    session: &mut dyn tenferro_tensor::BackendSession,
463    policy: CpuBatchPolicy,
464    f: impl FnOnce(&mut dyn tenferro_tensor::BackendSession) -> R,
465) -> tenferro_tensor::Result<R> {
466    let Some(previous) =
467        crate::with_cpu_exec_session(session, |cpu| cpu.replace_batch_policy(policy))
468    else {
469        return Err(tenferro_tensor::Error::unsupported(
470            "tenferro_cpu::with_batch_policy",
471            "the session is not a CPU execution session and does not forward one; CPU batch \
472             policies apply only to CpuBackend sessions",
473        ));
474    };
475    // `f` runs on the caller's own session, so a wrapping session keeps its
476    // overrides inside the scope. The policy is restored before an unwind
477    // continues; nothing observes the session between the panic and the
478    // restore, so asserting unwind safety is sound.
479    let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| f(&mut *session)));
480    // INVARIANT: the same session yielded a CPU execution session above and
481    // native tokens are stable for a session's lifetime, so this visit runs.
482    let _ = crate::with_cpu_exec_session(session, |cpu| cpu.replace_batch_policy(previous));
483    match outcome {
484        Ok(value) => Ok(value),
485        Err(payload) => std::panic::resume_unwind(payload),
486    }
487}
488
489/// Typed error for a forced batch strategy with no route for this operation.
490pub(crate) fn strategy_unavailable(
491    op: &'static str,
492    strategy: CpuBatchStrategy,
493    reason: &str,
494) -> tenferro_tensor::Error {
495    tenferro_tensor::Error::unsupported(
496        op,
497        format!(
498            "batch strategy {strategy:?} is not available: {reason}; use CpuBatchStrategy::Auto or another strategy"
499        ),
500    )
501}