use axiolid_core::{Scalar, Tolerance};
use crate::cancel::CancellationToken;
use crate::{BackendId, GeomError, GeomResult, Precision};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum Determinism {
BestEffort,
Topological,
NumericallyBounded,
Bitwise,
}
impl Determinism {
pub const fn satisfies(self, required: Self) -> bool {
(self as u8) >= (required as u8)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Parallelism {
Serial,
Auto,
Threads(usize),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum DevicePreference {
Auto,
Cpu,
Gpu,
Backend(BackendId),
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Residency {
Host,
Device(BackendId),
Unified(BackendId),
}
impl Residency {
pub fn is_local_to(self, backend: BackendId) -> bool {
match self {
Self::Host => false,
Self::Device(owner) | Self::Unified(owner) => owner == backend,
}
}
pub const fn is_host_readable(self) -> bool {
matches!(self, Self::Host | Self::Unified(_))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct DataResidency {
input: Residency,
output: Residency,
}
impl DataResidency {
pub const HOST: Self = Self {
input: Residency::Host,
output: Residency::Host,
};
pub const fn new(input: Residency, output: Residency) -> Self {
Self { input, output }
}
pub const fn input(self) -> Residency {
self.input
}
pub const fn output(self) -> Residency {
self.output
}
pub fn is_transfer_free_on(self, backend: BackendId) -> bool {
self.input.is_local_to(backend) && self.output.is_local_to(backend)
}
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ScratchRequirement {
None,
Fixed {
bytes: usize,
},
PerElement {
bytes_per_element: usize,
},
Unbounded,
}
impl ScratchRequirement {
pub const fn upper_bound_bytes(self, elements: usize) -> Option<usize> {
match self {
Self::None => Some(0),
Self::Fixed { bytes } => Some(bytes),
Self::PerElement { bytes_per_element } => bytes_per_element.checked_mul(elements),
Self::Unbounded => None,
}
}
pub fn fits_budget(self, options: &ExecutionOptions, elements: usize) -> bool {
match options.memory_budget_bytes() {
None => true,
Some(budget) => match self.upper_bound_bytes(elements) {
Some(needed) => needed <= budget,
None => false,
},
}
}
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum OutputBound {
OneToOne,
AtMost {
max: usize,
},
Unbounded,
}
impl OutputBound {
pub const fn upper_bound(self, elements: usize) -> Option<usize> {
match self {
Self::OneToOne => Some(elements),
Self::AtMost { max } => max.checked_mul(elements),
Self::Unbounded => None,
}
}
pub const fn is_preallocatable(self, elements: usize) -> bool {
self.upper_bound(elements).is_some()
}
pub fn write_offsets(self, counts: &[usize]) -> Option<(Vec<usize>, usize)> {
let per_element_max = match self {
Self::OneToOne => Some(1),
Self::AtMost { max } => Some(max),
Self::Unbounded => None,
};
let mut offsets = Vec::with_capacity(counts.len());
let mut running = 0usize;
for &count in counts {
if let Some(max) = per_element_max {
if count > max {
return None;
}
}
offsets.push(running);
running = running.checked_add(count)?;
}
Some((offsets, running))
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ExecutionOptions {
tolerance: Tolerance,
precision: Precision,
determinism: Determinism,
parallelism: Parallelism,
device: DevicePreference,
residency: DataResidency,
memory_budget_bytes: Option<usize>,
cancellation: Option<CancellationToken>,
chord_error: Option<Scalar>,
}
impl ExecutionOptions {
pub const fn new(tolerance: Tolerance) -> Self {
Self {
tolerance,
precision: Precision::F64,
determinism: Determinism::NumericallyBounded,
parallelism: Parallelism::Auto,
device: DevicePreference::Auto,
residency: DataResidency::HOST,
memory_budget_bytes: None,
cancellation: None,
chord_error: None,
}
}
pub fn with_tolerance(mut self, tolerance: Tolerance) -> Self {
self.tolerance = tolerance;
self
}
pub fn with_cancellation(mut self, token: CancellationToken) -> Self {
self.cancellation = Some(token);
self
}
pub fn cancellation(&self) -> Option<&CancellationToken> {
self.cancellation.as_ref()
}
pub fn check_cancelled(&self) -> crate::GeomResult<()> {
match &self.cancellation {
Some(token) => token.check(),
None => Ok(()),
}
}
pub fn with_precision(mut self, precision: Precision) -> Self {
self.precision = precision;
self
}
pub fn with_determinism(mut self, value: Determinism) -> Self {
self.determinism = value;
self
}
pub fn with_parallelism(mut self, value: Parallelism) -> Option<Self> {
if matches!(value, Parallelism::Threads(0)) {
return None;
}
self.parallelism = value;
Some(self)
}
pub fn with_device(mut self, value: DevicePreference) -> Self {
self.device = value;
self
}
pub fn with_residency(mut self, value: DataResidency) -> Self {
self.residency = value;
self
}
pub fn with_chord_error(mut self, value: Scalar) -> Option<Self> {
if !(value.is_finite() && value > 0.0) {
return None;
}
self.chord_error = Some(value);
Some(self)
}
pub fn chord_error(&self) -> Option<Scalar> {
self.chord_error
}
pub fn with_memory_budget(mut self, bytes: usize) -> Self {
self.memory_budget_bytes = Some(bytes);
self
}
pub fn tolerance(&self) -> Tolerance {
self.tolerance
}
pub fn precision(&self) -> Precision {
self.precision
}
pub fn determinism(&self) -> Determinism {
self.determinism
}
pub fn parallelism(&self) -> Parallelism {
self.parallelism
}
pub fn device(&self) -> DevicePreference {
self.device
}
pub fn memory_budget_bytes(&self) -> Option<usize> {
self.memory_budget_bytes
}
pub fn residency(&self) -> DataResidency {
self.residency
}
pub fn charge_scratch(&self, bytes: usize) -> GeomResult<()> {
match self.memory_budget_bytes {
Some(budget) if bytes > budget => Err(GeomError::BudgetExceeded { resource: "memory" }),
_ => Ok(()),
}
}
}