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}