use std::collections::BTreeMap;
use smallvec::SmallVec;
use tenferro_tensor::{CpuDomainId, DType, Tensor};
const INLINE_DOMAIN_CAPACITY: usize = 8;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum CpuAffinityPolicy {
DominantInputBytes,
RequireSingleDomain,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct CpuAffinityInput {
pub domain: Option<CpuDomainId>,
pub logical_bytes: usize,
}
impl CpuAffinityInput {
pub fn from_tensor(tensor: &Tensor) -> Result<Self, CpuAffinityInputError> {
Self::from_parts(
tensor.placement().cpu_affinity,
tensor.shape(),
tensor.dtype(),
)
}
pub fn from_parts(
domain: Option<CpuDomainId>,
shape: &[usize],
dtype: DType,
) -> Result<Self, CpuAffinityInputError> {
let element_count = if shape.contains(&0) {
0
} else {
shape.iter().try_fold(1_usize, |count, &extent| {
count
.checked_mul(extent)
.ok_or(CpuAffinityInputError::ShapeProductOverflow)
})?
};
let byte_width = dtype_byte_width(dtype);
let logical_bytes = element_count.checked_mul(byte_width).ok_or(
CpuAffinityInputError::LogicalByteCountOverflow {
element_count,
byte_width,
},
)?;
Ok(Self {
domain,
logical_bytes,
})
}
}
const fn dtype_byte_width(dtype: DType) -> usize {
match dtype {
DType::F32 | DType::I32 => std::mem::size_of::<u32>(),
DType::F64 | DType::I64 => std::mem::size_of::<u64>(),
DType::Bool => std::mem::size_of::<bool>(),
DType::C32 => std::mem::size_of::<num_complex::Complex32>(),
DType::C64 => std::mem::size_of::<num_complex::Complex64>(),
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
pub enum CpuAffinityInputError {
#[error("logical tensor element count overflowed usize")]
ShapeProductOverflow,
#[error(
"logical tensor byte count overflowed: element_count={element_count}, byte_width={byte_width}"
)]
LogicalByteCountOverflow {
element_count: usize,
byte_width: usize,
},
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum CpuAffinitySelectionReason {
ExplicitOverride,
DominantInputBytes,
SingleInputDomain,
DefaultDomain,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct CpuAffinitySelection {
pub domain: CpuDomainId,
pub reason: CpuAffinitySelectionReason,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
pub enum CpuAffinityResolutionError {
#[error("logical input-byte total overflowed for CPU domain {domain:?}")]
LogicalByteCountOverflow {
domain: CpuDomainId,
},
#[error("CPU affinity policy requires one input domain, found {first:?} and {second:?}")]
MultipleKnownDomains {
first: CpuDomainId,
second: CpuDomainId,
},
}
pub fn resolve_cpu_affinity(
policy: CpuAffinityPolicy,
inputs: &[CpuAffinityInput],
default_domain: CpuDomainId,
) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
resolve_cpu_affinity_with_override(policy, inputs, default_domain, None)
}
pub fn resolve_cpu_affinity_with_override(
policy: CpuAffinityPolicy,
inputs: &[CpuAffinityInput],
default_domain: CpuDomainId,
explicit_domain: Option<CpuDomainId>,
) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
if let Some(domain) = explicit_domain {
return Ok(CpuAffinitySelection {
domain,
reason: CpuAffinitySelectionReason::ExplicitOverride,
});
}
match policy {
CpuAffinityPolicy::DominantInputBytes => resolve_dominant(inputs, default_domain),
CpuAffinityPolicy::RequireSingleDomain => resolve_single(inputs, default_domain),
}
}
fn resolve_dominant(
inputs: &[CpuAffinityInput],
default_domain: CpuDomainId,
) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
let mut totals = DomainTotals::default();
for input in inputs {
let Some(domain) = input.domain else {
continue;
};
if input.logical_bytes == 0 {
continue;
}
totals.add(domain, input.logical_bytes);
}
if let Some(domain) = totals.smallest_overflowing_domain() {
return Err(CpuAffinityResolutionError::LogicalByteCountOverflow { domain });
}
match totals.dominant_domain() {
Some(domain) => Ok(CpuAffinitySelection {
domain,
reason: CpuAffinitySelectionReason::DominantInputBytes,
}),
None => Ok(default_selection(default_domain)),
}
}
fn resolve_single(
inputs: &[CpuAffinityInput],
default_domain: CpuDomainId,
) -> Result<CpuAffinitySelection, CpuAffinityResolutionError> {
let mut first = None;
let mut second = None;
for domain in inputs.iter().filter_map(|input| input.domain) {
observe_smallest_two_distinct(domain, &mut first, &mut second);
}
if let (Some(first), Some(second)) = (first, second) {
return Err(CpuAffinityResolutionError::MultipleKnownDomains { first, second });
}
Ok(match first {
Some(domain) => CpuAffinitySelection {
domain,
reason: CpuAffinitySelectionReason::SingleInputDomain,
},
None => default_selection(default_domain),
})
}
fn observe_smallest_two_distinct(
domain: CpuDomainId,
first: &mut Option<CpuDomainId>,
second: &mut Option<CpuDomainId>,
) {
if *first == Some(domain) || *second == Some(domain) {
return;
}
match *first {
None => *first = Some(domain),
Some(current_first) if domain < current_first => {
*second = *first;
*first = Some(domain);
}
Some(_) if second.is_none_or(|current_second| domain < current_second) => {
*second = Some(domain);
}
Some(_) => {}
}
}
fn default_selection(domain: CpuDomainId) -> CpuAffinitySelection {
CpuAffinitySelection {
domain,
reason: CpuAffinitySelectionReason::DefaultDomain,
}
}
#[derive(Clone, Copy, Debug)]
struct DomainTotal {
domain: CpuDomainId,
logical_bytes: Option<usize>,
}
impl DomainTotal {
fn new(domain: CpuDomainId, logical_bytes: usize) -> Self {
Self {
domain,
logical_bytes: Some(logical_bytes),
}
}
fn add(&mut self, logical_bytes: usize) {
self.logical_bytes = self
.logical_bytes
.and_then(|total| total.checked_add(logical_bytes));
}
}
enum DomainTotals {
Inline(SmallVec<[DomainTotal; INLINE_DOMAIN_CAPACITY]>),
Heap(BTreeMap<CpuDomainId, Option<usize>>),
}
impl Default for DomainTotals {
fn default() -> Self {
Self::Inline(SmallVec::new())
}
}
impl DomainTotals {
fn add(&mut self, domain: CpuDomainId, logical_bytes: usize) {
let promoted = match self {
Self::Inline(entries) => {
if let Some(entry) = entries.iter_mut().find(|entry| entry.domain == domain) {
entry.add(logical_bytes);
return;
}
if entries.len() < INLINE_DOMAIN_CAPACITY {
entries.push(DomainTotal::new(domain, logical_bytes));
return;
}
let mut heap = BTreeMap::new();
for entry in entries.drain(..) {
heap.insert(entry.domain, entry.logical_bytes);
}
heap.insert(domain, Some(logical_bytes));
Some(heap)
}
Self::Heap(entries) => {
let total = entries.entry(domain).or_insert(Some(0));
*total = total.and_then(|current| current.checked_add(logical_bytes));
None
}
};
if let Some(heap) = promoted {
*self = Self::Heap(heap);
}
}
fn smallest_overflowing_domain(&self) -> Option<CpuDomainId> {
match self {
Self::Inline(entries) => entries
.iter()
.filter(|entry| entry.logical_bytes.is_none())
.map(|entry| entry.domain)
.min(),
Self::Heap(entries) => entries
.iter()
.find_map(|(domain, total)| total.is_none().then_some(*domain)),
}
}
fn dominant_domain(&self) -> Option<CpuDomainId> {
let mut best = None;
match self {
Self::Inline(entries) => {
for entry in entries {
if let Some(logical_bytes) = entry.logical_bytes {
consider_dominant(&mut best, entry.domain, logical_bytes);
}
}
}
Self::Heap(entries) => {
for (&domain, &logical_bytes) in entries {
if let Some(logical_bytes) = logical_bytes {
consider_dominant(&mut best, domain, logical_bytes);
}
}
}
}
best.map(|(domain, _)| domain)
}
}
fn consider_dominant(
best: &mut Option<(CpuDomainId, usize)>,
domain: CpuDomainId,
logical_bytes: usize,
) {
let replace = match *best {
None => true,
Some((best_domain, best_bytes)) => {
logical_bytes > best_bytes || (logical_bytes == best_bytes && domain < best_domain)
}
};
if replace {
*best = Some((domain, logical_bytes));
}
}
#[cfg(test)]
mod tests;