tenferro_cpu/affinity_policy.rs
1use std::collections::BTreeMap;
2
3use smallvec::SmallVec;
4use tenferro_tensor::{CpuDomainId, DType, Tensor};
5
6const INLINE_DOMAIN_CAPACITY: usize = 8;
7
8/// Policy used to select a CPU execution domain from input affinity metadata.
9///
10/// # Examples
11///
12/// ```rust
13/// use tenferro_cpu::CpuAffinityPolicy;
14///
15/// let policy = CpuAffinityPolicy::DominantInputBytes;
16/// assert_ne!(policy, CpuAffinityPolicy::RequireSingleDomain);
17/// ```
18#[derive(Clone, Copy, Debug, Eq, PartialEq)]
19pub enum CpuAffinityPolicy {
20 /// Select the domain with the largest total of positive logical input bytes.
21 DominantInputBytes,
22 /// Accept zero or one known input domain and reject mixed known domains.
23 RequireSingleDomain,
24}
25
26/// CPU affinity metadata for one logical operation input.
27///
28/// The resolver reads this metadata only. It never changes, copies, or rehomes
29/// tensor payloads.
30///
31/// # Examples
32///
33/// ```rust
34/// use tenferro_cpu::CpuAffinityInput;
35/// use tenferro_tensor::CpuDomainId;
36///
37/// let input = CpuAffinityInput {
38/// domain: Some(CpuDomainId::new(3)),
39/// logical_bytes: 64,
40/// };
41/// assert_eq!(input.logical_bytes, 64);
42/// ```
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub struct CpuAffinityInput {
45 /// Known CPU execution domain, or `None` when affinity is unknown.
46 pub domain: Option<CpuDomainId>,
47 /// Logical input size used by [`CpuAffinityPolicy::DominantInputBytes`].
48 pub logical_bytes: usize,
49}
50
51impl CpuAffinityInput {
52 /// Construct resolver input metadata from a tensor.
53 ///
54 /// The logical byte count is the checked shape product times the tensor's
55 /// scalar byte width. CPU affinity is copied from placement metadata; the
56 /// tensor and its storage are otherwise untouched.
57 ///
58 /// # Examples
59 ///
60 /// ```rust
61 /// use tenferro_cpu::CpuAffinityInput;
62 /// use tenferro_tensor::Tensor;
63 ///
64 /// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
65 /// let input = CpuAffinityInput::from_tensor(&tensor)?;
66 /// assert_eq!(input.logical_bytes, 2 * std::mem::size_of::<f64>());
67 /// # Ok::<(), Box<dyn std::error::Error>>(())
68 /// ```
69 ///
70 /// # Errors
71 ///
72 /// Returns [`CpuAffinityInputError`] when the logical element or byte
73 /// count cannot be represented by `usize`.
74 pub fn from_tensor(tensor: &Tensor) -> Result<Self, CpuAffinityInputError> {
75 Self::from_parts(
76 tensor.placement().cpu_affinity,
77 tensor.shape(),
78 tensor.dtype(),
79 )
80 }
81
82 /// Construct resolver input metadata from placement, shape, and dtype.
83 ///
84 /// Scalar shapes have one element. Any zero extent yields zero logical
85 /// bytes without multiplying the other extents.
86 ///
87 /// # Examples
88 ///
89 /// ```rust
90 /// use tenferro_cpu::CpuAffinityInput;
91 /// use tenferro_tensor::DType;
92 ///
93 /// let scalar = CpuAffinityInput::from_parts(None, &[], DType::F32)?;
94 /// let empty = CpuAffinityInput::from_parts(None, &[usize::MAX, 0], DType::F64)?;
95 /// assert_eq!(scalar.logical_bytes, 4);
96 /// assert_eq!(empty.logical_bytes, 0);
97 /// # Ok::<(), tenferro_cpu::CpuAffinityInputError>(())
98 /// ```
99 ///
100 /// # Errors
101 ///
102 /// Returns [`CpuAffinityInputError::ShapeProductOverflow`] when non-zero
103 /// extents overflow, or
104 /// [`CpuAffinityInputError::LogicalByteCountOverflow`] when multiplying by
105 /// the dtype width overflows.
106 pub fn from_parts(
107 domain: Option<CpuDomainId>,
108 shape: &[usize],
109 dtype: DType,
110 ) -> Result<Self, CpuAffinityInputError> {
111 let element_count = if shape.contains(&0) {
112 0
113 } else {
114 shape.iter().try_fold(1_usize, |count, &extent| {
115 count
116 .checked_mul(extent)
117 .ok_or(CpuAffinityInputError::ShapeProductOverflow)
118 })?
119 };
120 let byte_width = dtype_byte_width(dtype);
121 let logical_bytes = element_count.checked_mul(byte_width).ok_or(
122 CpuAffinityInputError::LogicalByteCountOverflow {
123 element_count,
124 byte_width,
125 },
126 )?;
127 Ok(Self {
128 domain,
129 logical_bytes,
130 })
131 }
132}
133
134const fn dtype_byte_width(dtype: DType) -> usize {
135 match dtype {
136 DType::F32 | DType::I32 => std::mem::size_of::<u32>(),
137 DType::F64 | DType::I64 => std::mem::size_of::<u64>(),
138 DType::Bool => std::mem::size_of::<bool>(),
139 DType::C32 => std::mem::size_of::<num_complex::Complex32>(),
140 DType::C64 => std::mem::size_of::<num_complex::Complex64>(),
141 // INVARIANT: CPU affinity is derived for admitted preset tensors, whose
142 // element width is fixed. An externally defined scalar is caller-owned and
143 // has no fixed width here.
144 DType::External(_) => 0,
145 }
146}
147
148/// Failure to derive logical input bytes for CPU affinity resolution.
149///
150/// # Examples
151///
152/// ```rust
153/// use tenferro_cpu::{CpuAffinityInput, CpuAffinityInputError};
154/// use tenferro_tensor::DType;
155///
156/// let error = CpuAffinityInput::from_parts(None, &[usize::MAX, 2], DType::F32)
157/// .unwrap_err();
158/// assert_eq!(error, CpuAffinityInputError::ShapeProductOverflow);
159/// ```
160#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
161pub enum CpuAffinityInputError {
162 /// Multiplying non-zero shape extents overflowed `usize`.
163 #[error("logical tensor element count overflowed usize")]
164 ShapeProductOverflow,
165 /// Multiplying element count by dtype width overflowed `usize`.
166 #[error(
167 "logical tensor byte count overflowed: element_count={element_count}, byte_width={byte_width}"
168 )]
169 LogicalByteCountOverflow {
170 /// Checked logical element count.
171 element_count: usize,
172 /// Scalar dtype width in bytes.
173 byte_width: usize,
174 },
175}
176
177/// Why the CPU affinity resolver selected a domain.
178///
179/// # Examples
180///
181/// ```rust
182/// use tenferro_cpu::CpuAffinitySelectionReason;
183///
184/// let reason = CpuAffinitySelectionReason::DefaultDomain;
185/// assert_eq!(reason, CpuAffinitySelectionReason::DefaultDomain);
186/// ```
187#[derive(Clone, Copy, Debug, Eq, PartialEq)]
188pub enum CpuAffinitySelectionReason {
189 /// An operation-local explicit domain override took precedence.
190 ExplicitOverride,
191 /// Positive logical bytes made this domain dominant.
192 DominantInputBytes,
193 /// Strict policy observed exactly one known input domain.
194 SingleInputDomain,
195 /// No relevant input affinity was available.
196 DefaultDomain,
197}
198
199/// Deterministic CPU affinity selection returned by the pure resolver.
200///
201/// # Examples
202///
203/// ```rust
204/// use tenferro_cpu::{CpuAffinitySelection, CpuAffinitySelectionReason};
205/// use tenferro_tensor::CpuDomainId;
206///
207/// let selection = CpuAffinitySelection {
208/// domain: CpuDomainId::new(2),
209/// reason: CpuAffinitySelectionReason::DominantInputBytes,
210/// };
211/// assert_eq!(selection.domain.as_u64(), 2);
212/// ```
213#[derive(Clone, Copy, Debug, Eq, PartialEq)]
214pub struct CpuAffinitySelection {
215 /// Selected CPU execution domain.
216 pub domain: CpuDomainId,
217 /// Deterministic reason for the selection.
218 pub reason: CpuAffinitySelectionReason,
219}
220
221/// Failure to resolve CPU affinity from input metadata.
222///
223/// # Examples
224///
225/// ```rust
226/// use tenferro_cpu::CpuAffinityResolutionError;
227/// use tenferro_tensor::CpuDomainId;
228///
229/// let error = CpuAffinityResolutionError::LogicalByteCountOverflow {
230/// domain: CpuDomainId::new(4),
231/// };
232/// assert!(error.to_string().contains("4"));
233/// ```
234#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
235pub enum CpuAffinityResolutionError {
236 /// Adding logical byte counts overflowed `usize` for one domain.
237 #[error("logical input-byte total overflowed for CPU domain {domain:?}")]
238 LogicalByteCountOverflow {
239 /// Smallest CPU domain whose logical byte total overflowed.
240 domain: CpuDomainId,
241 },
242 /// Strict policy observed at least two different known domains.
243 #[error("CPU affinity policy requires one input domain, found {first:?} and {second:?}")]
244 MultipleKnownDomains {
245 /// Smallest known input domain.
246 first: CpuDomainId,
247 /// Second-smallest known input domain.
248 second: CpuDomainId,
249 },
250}
251
252/// Resolve a CPU execution domain from input affinity metadata.
253///
254/// Unknown affinities and zero-byte inputs do not contribute to dominant-byte
255/// scoring. When no input contributes, `default_domain` is selected. Equal
256/// positive totals are resolved in favor of the smallest [`CpuDomainId`]. The
257/// input slice is only read; the resolver never retags or rehomes an input.
258///
259/// Use [`resolve_cpu_affinity_with_override`] when an operation-local explicit
260/// placement has already been selected.
261///
262/// # Examples
263///
264/// ```rust
265/// use tenferro_cpu::{resolve_cpu_affinity, CpuAffinityInput, CpuAffinityPolicy};
266/// use tenferro_tensor::CpuDomainId;
267///
268/// let inputs = [
269/// CpuAffinityInput { domain: Some(CpuDomainId::new(8)), logical_bytes: 6 },
270/// CpuAffinityInput { domain: Some(CpuDomainId::new(3)), logical_bytes: 2 },
271/// ];
272/// let selected = resolve_cpu_affinity(
273/// CpuAffinityPolicy::DominantInputBytes,
274/// &inputs,
275/// CpuDomainId::new(1),
276/// )?;
277/// assert_eq!(selected.domain, CpuDomainId::new(8));
278/// # Ok::<(), tenferro_cpu::CpuAffinityResolutionError>(())
279/// ```
280///
281/// # Errors
282///
283/// Returns [`CpuAffinityResolutionError::LogicalByteCountOverflow`] when one
284/// domain's logical byte total cannot be represented by `usize`, or
285/// [`CpuAffinityResolutionError::MultipleKnownDomains`] when strict policy sees
286/// more than one known input domain.
287pub fn resolve_cpu_affinity(
288 policy: CpuAffinityPolicy,
289 inputs: &[CpuAffinityInput],
290 default_domain: CpuDomainId,
291) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
292 resolve_cpu_affinity_with_override(policy, inputs, default_domain, None)
293}
294
295/// Resolve CPU affinity with an optional operation-local explicit override.
296///
297/// Explicit placement takes precedence before input-byte accounting or strict
298/// mixed-domain validation. Passing `None` applies the same policy resolution
299/// as [`resolve_cpu_affinity`].
300///
301/// # Examples
302///
303/// ```rust
304/// use tenferro_cpu::{
305/// resolve_cpu_affinity_with_override, CpuAffinityInput, CpuAffinityPolicy,
306/// CpuAffinitySelectionReason,
307/// };
308/// use tenferro_tensor::CpuDomainId;
309///
310/// let mixed = [
311/// CpuAffinityInput { domain: Some(CpuDomainId::new(1)), logical_bytes: 1 },
312/// CpuAffinityInput { domain: Some(CpuDomainId::new(2)), logical_bytes: 1 },
313/// ];
314/// let selected = resolve_cpu_affinity_with_override(
315/// CpuAffinityPolicy::RequireSingleDomain,
316/// &mixed,
317/// CpuDomainId::new(1),
318/// Some(CpuDomainId::new(9)),
319/// )?;
320/// assert_eq!(selected.domain, CpuDomainId::new(9));
321/// assert_eq!(selected.reason, CpuAffinitySelectionReason::ExplicitOverride);
322/// # Ok::<(), tenferro_cpu::CpuAffinityResolutionError>(())
323/// ```
324///
325/// # Errors
326///
327/// When `explicit_domain` is `None`, returns
328/// [`CpuAffinityResolutionError::LogicalByteCountOverflow`] for an unrepresentable
329/// domain byte total or [`CpuAffinityResolutionError::MultipleKnownDomains`]
330/// when strict policy sees more than one known input domain. A present explicit
331/// override bypasses both policy errors.
332pub fn resolve_cpu_affinity_with_override(
333 policy: CpuAffinityPolicy,
334 inputs: &[CpuAffinityInput],
335 default_domain: CpuDomainId,
336 explicit_domain: Option<CpuDomainId>,
337) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
338 if let Some(domain) = explicit_domain {
339 return Ok(CpuAffinitySelection {
340 domain,
341 reason: CpuAffinitySelectionReason::ExplicitOverride,
342 });
343 }
344 match policy {
345 CpuAffinityPolicy::DominantInputBytes => resolve_dominant(inputs, default_domain),
346 CpuAffinityPolicy::RequireSingleDomain => resolve_single(inputs, default_domain),
347 }
348}
349
350fn resolve_dominant(
351 inputs: &[CpuAffinityInput],
352 default_domain: CpuDomainId,
353) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
354 let mut totals = DomainTotals::default();
355 for input in inputs {
356 let Some(domain) = input.domain else {
357 continue;
358 };
359 if input.logical_bytes == 0 {
360 continue;
361 }
362 totals.add(domain, input.logical_bytes);
363 }
364
365 if let Some(domain) = totals.smallest_overflowing_domain() {
366 return Err(CpuAffinityResolutionError::LogicalByteCountOverflow { domain });
367 }
368
369 match totals.dominant_domain() {
370 Some(domain) => Ok(CpuAffinitySelection {
371 domain,
372 reason: CpuAffinitySelectionReason::DominantInputBytes,
373 }),
374 None => Ok(default_selection(default_domain)),
375 }
376}
377
378fn resolve_single(
379 inputs: &[CpuAffinityInput],
380 default_domain: CpuDomainId,
381) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
382 let mut first = None;
383 let mut second = None;
384 for domain in inputs.iter().filter_map(|input| input.domain) {
385 observe_smallest_two_distinct(domain, &mut first, &mut second);
386 }
387
388 if let (Some(first), Some(second)) = (first, second) {
389 return Err(CpuAffinityResolutionError::MultipleKnownDomains { first, second });
390 }
391
392 Ok(match first {
393 Some(domain) => CpuAffinitySelection {
394 domain,
395 reason: CpuAffinitySelectionReason::SingleInputDomain,
396 },
397 None => default_selection(default_domain),
398 })
399}
400
401fn observe_smallest_two_distinct(
402 domain: CpuDomainId,
403 first: &mut Option<CpuDomainId>,
404 second: &mut Option<CpuDomainId>,
405) {
406 if *first == Some(domain) || *second == Some(domain) {
407 return;
408 }
409 match *first {
410 None => *first = Some(domain),
411 Some(current_first) if domain < current_first => {
412 *second = *first;
413 *first = Some(domain);
414 }
415 Some(_) if second.is_none_or(|current_second| domain < current_second) => {
416 *second = Some(domain);
417 }
418 Some(_) => {}
419 }
420}
421
422fn default_selection(domain: CpuDomainId) -> CpuAffinitySelection {
423 CpuAffinitySelection {
424 domain,
425 reason: CpuAffinitySelectionReason::DefaultDomain,
426 }
427}
428
429#[derive(Clone, Copy, Debug)]
430struct DomainTotal {
431 domain: CpuDomainId,
432 logical_bytes: Option<usize>,
433}
434
435impl DomainTotal {
436 fn new(domain: CpuDomainId, logical_bytes: usize) -> Self {
437 Self {
438 domain,
439 logical_bytes: Some(logical_bytes),
440 }
441 }
442
443 fn add(&mut self, logical_bytes: usize) {
444 self.logical_bytes = self
445 .logical_bytes
446 .and_then(|total| total.checked_add(logical_bytes));
447 }
448}
449
450enum DomainTotals {
451 Inline(SmallVec<[DomainTotal; INLINE_DOMAIN_CAPACITY]>),
452 Heap(BTreeMap<CpuDomainId, Option<usize>>),
453}
454
455impl Default for DomainTotals {
456 fn default() -> Self {
457 Self::Inline(SmallVec::new())
458 }
459}
460
461impl DomainTotals {
462 fn add(&mut self, domain: CpuDomainId, logical_bytes: usize) {
463 let promoted = match self {
464 Self::Inline(entries) => {
465 // INVARIANT: the linear lookup is bounded by the inline capacity;
466 // larger distinct-domain sets are promoted to `BTreeMap` below.
467 if let Some(entry) = entries.iter_mut().find(|entry| entry.domain == domain) {
468 entry.add(logical_bytes);
469 return;
470 }
471 if entries.len() < INLINE_DOMAIN_CAPACITY {
472 entries.push(DomainTotal::new(domain, logical_bytes));
473 return;
474 }
475 let mut heap = BTreeMap::new();
476 for entry in entries.drain(..) {
477 heap.insert(entry.domain, entry.logical_bytes);
478 }
479 heap.insert(domain, Some(logical_bytes));
480 Some(heap)
481 }
482 Self::Heap(entries) => {
483 let total = entries.entry(domain).or_insert(Some(0));
484 *total = total.and_then(|current| current.checked_add(logical_bytes));
485 None
486 }
487 };
488 if let Some(heap) = promoted {
489 *self = Self::Heap(heap);
490 }
491 }
492
493 fn smallest_overflowing_domain(&self) -> Option<CpuDomainId> {
494 match self {
495 Self::Inline(entries) => entries
496 .iter()
497 .filter(|entry| entry.logical_bytes.is_none())
498 .map(|entry| entry.domain)
499 .min(),
500 Self::Heap(entries) => entries
501 .iter()
502 .find_map(|(domain, total)| total.is_none().then_some(*domain)),
503 }
504 }
505
506 fn dominant_domain(&self) -> Option<CpuDomainId> {
507 let mut best = None;
508 match self {
509 Self::Inline(entries) => {
510 for entry in entries {
511 if let Some(logical_bytes) = entry.logical_bytes {
512 consider_dominant(&mut best, entry.domain, logical_bytes);
513 }
514 }
515 }
516 Self::Heap(entries) => {
517 for (&domain, &logical_bytes) in entries {
518 if let Some(logical_bytes) = logical_bytes {
519 consider_dominant(&mut best, domain, logical_bytes);
520 }
521 }
522 }
523 }
524 best.map(|(domain, _)| domain)
525 }
526}
527
528fn consider_dominant(
529 best: &mut Option<(CpuDomainId, usize)>,
530 domain: CpuDomainId,
531 logical_bytes: usize,
532) {
533 let replace = match *best {
534 None => true,
535 Some((best_domain, best_bytes)) => {
536 logical_bytes > best_bytes || (logical_bytes == best_bytes && domain < best_domain)
537 }
538 };
539 if replace {
540 *best = Some((domain, logical_bytes));
541 }
542}
543
544#[cfg(test)]
545mod tests;