Skip to main content

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;