Skip to main content

runmat_accelerate_api/placement/
representation.rs

1use serde::{Deserialize, Serialize};
2
3use crate::{IntegerElementType, ProviderPrecision};
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
6#[serde(rename_all = "snake_case")]
7pub enum ProviderElementType {
8    Logical,
9    I8,
10    I16,
11    I32,
12    I64,
13    U8,
14    U16,
15    U32,
16    U64,
17    F32,
18    F64,
19    ComplexF32,
20    ComplexF64,
21}
22
23impl From<ProviderPrecision> for ProviderElementType {
24    fn from(value: ProviderPrecision) -> Self {
25        match value {
26            ProviderPrecision::F32 => Self::F32,
27            ProviderPrecision::F64 => Self::F64,
28        }
29    }
30}
31
32impl From<IntegerElementType> for ProviderElementType {
33    fn from(value: IntegerElementType) -> Self {
34        match value {
35            IntegerElementType::I8 => Self::I8,
36            IntegerElementType::I16 => Self::I16,
37            IntegerElementType::I32 => Self::I32,
38            IntegerElementType::I64 => Self::I64,
39            IntegerElementType::U8 => Self::U8,
40            IntegerElementType::U16 => Self::U16,
41            IntegerElementType::U32 => Self::U32,
42            IntegerElementType::U64 => Self::U64,
43        }
44    }
45}
46
47impl ProviderElementType {
48    pub const fn byte_width(self) -> u64 {
49        match self {
50            Self::Logical | Self::I8 | Self::U8 => 1,
51            Self::I16 | Self::U16 => 2,
52            Self::I32 | Self::U32 | Self::F32 => 4,
53            Self::I64 | Self::U64 | Self::F64 => 8,
54            Self::ComplexF32 => 8,
55            Self::ComplexF64 => 16,
56        }
57    }
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
61#[serde(rename_all = "snake_case")]
62pub enum ProviderStorage {
63    DenseReal,
64    DenseComplexInterleaved,
65    Sparse,
66}
67
68#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
69#[serde(rename_all = "snake_case")]
70pub enum ProviderLayout {
71    ColumnMajorContiguous,
72    Strided,
73    Opaque,
74}
75
76#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
77#[serde(rename_all = "snake_case")]
78pub enum ProviderResidency {
79    Host,
80    Device,
81    Mirrored,
82    Unknown,
83}
84
85#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
86#[serde(deny_unknown_fields)]
87pub struct ProviderRepresentation {
88    pub element_type: ProviderElementType,
89    pub storage: ProviderStorage,
90    pub layout: ProviderLayout,
91    pub shape: Vec<u64>,
92    pub residency: ProviderResidency,
93}
94
95impl ProviderRepresentation {
96    pub fn checked_element_count(&self) -> Option<u64> {
97        self.shape
98            .iter()
99            .try_fold(1_u64, |count, extent| count.checked_mul(*extent))
100    }
101
102    pub fn checked_byte_len(&self) -> Option<u64> {
103        self.checked_element_count()?
104            .checked_mul(self.element_type.byte_width())
105    }
106}