Skip to main content

rlmesh_spaces/
dtype.rs

1/// Element data type for tensor values exchanged across the wire.
2///
3/// Discriminants are kept byte-identical to the `rlmesh.spaces.v1.DType` proto
4/// enum. The equality is enforced at compile time by a cross-check in
5/// `rlmesh-grpc` (`wire::spaces`); `rlmesh-spaces` stays free of a
6/// `rlmesh-proto` dependency, so the assert cannot live here.
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
8#[repr(i32)]
9pub enum DType {
10    #[default]
11    Unspecified = 0,
12    Bool = 1,
13    Uint8 = 2,
14    Uint16 = 3,
15    Uint32 = 4,
16    Uint64 = 5,
17    Int8 = 6,
18    Int16 = 7,
19    Int32 = 8,
20    Int64 = 9,
21    Float16 = 10,
22    Float32 = 11,
23    Float64 = 12,
24}
25
26impl TryFrom<i32> for DType {
27    type Error = &'static str;
28
29    fn try_from(value: i32) -> Result<Self, Self::Error> {
30        match value {
31            0 => Ok(Self::Unspecified),
32            1 => Ok(Self::Bool),
33            2 => Ok(Self::Uint8),
34            3 => Ok(Self::Uint16),
35            4 => Ok(Self::Uint32),
36            5 => Ok(Self::Uint64),
37            6 => Ok(Self::Int8),
38            7 => Ok(Self::Int16),
39            8 => Ok(Self::Int32),
40            9 => Ok(Self::Int64),
41            10 => Ok(Self::Float16),
42            11 => Ok(Self::Float32),
43            12 => Ok(Self::Float64),
44            _ => Err("invalid dtype"),
45        }
46    }
47}
48
49impl From<DType> for i32 {
50    fn from(value: DType) -> Self {
51        value as i32
52    }
53}
54
55impl DType {
56    /// Every dtype, including `Unspecified`, in discriminant order.
57    pub const ALL: [DType; 13] = [
58        DType::Unspecified,
59        DType::Bool,
60        DType::Uint8,
61        DType::Uint16,
62        DType::Uint32,
63        DType::Uint64,
64        DType::Int8,
65        DType::Int16,
66        DType::Int32,
67        DType::Int64,
68        DType::Float16,
69        DType::Float32,
70        DType::Float64,
71    ];
72
73    /// The canonical lowercase dtype name (for example `"float32"`).
74    pub const fn name(self) -> &'static str {
75        match self {
76            DType::Unspecified => "unspecified",
77            DType::Bool => "bool",
78            DType::Uint8 => "uint8",
79            DType::Int32 => "int32",
80            DType::Int64 => "int64",
81            DType::Float16 => "float16",
82            DType::Float32 => "float32",
83            DType::Float64 => "float64",
84            DType::Int8 => "int8",
85            DType::Int16 => "int16",
86            DType::Uint16 => "uint16",
87            DType::Uint32 => "uint32",
88            DType::Uint64 => "uint64",
89        }
90    }
91
92    /// Parse a canonical dtype name. Only the 12 concrete dtypes are
93    /// recognized; `"unspecified"` and unknown names return `None`.
94    pub fn from_name(name: &str) -> Option<Self> {
95        match name {
96            "bool" => Some(DType::Bool),
97            "uint8" => Some(DType::Uint8),
98            "int32" => Some(DType::Int32),
99            "int64" => Some(DType::Int64),
100            "float16" => Some(DType::Float16),
101            "float32" => Some(DType::Float32),
102            "float64" => Some(DType::Float64),
103            "int8" => Some(DType::Int8),
104            "int16" => Some(DType::Int16),
105            "uint16" => Some(DType::Uint16),
106            "uint32" => Some(DType::Uint32),
107            "uint64" => Some(DType::Uint64),
108            _ => None,
109        }
110    }
111
112    /// Whether this is an integer dtype (signed or unsigned). `Bool` and the
113    /// float dtypes are not integers.
114    #[must_use]
115    pub const fn is_integer(self) -> bool {
116        matches!(
117            self,
118            DType::Uint8
119                | DType::Int8
120                | DType::Int16
121                | DType::Uint16
122                | DType::Int32
123                | DType::Uint32
124                | DType::Int64
125                | DType::Uint64
126        )
127    }
128
129    /// Whether this is a floating-point dtype.
130    #[must_use]
131    pub const fn is_float(self) -> bool {
132        matches!(self, DType::Float16 | DType::Float32 | DType::Float64)
133    }
134}
135
136impl std::fmt::Display for DType {
137    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
138        f.write_str(self.name())
139    }
140}
141
142/// Get the byte size of a dtype. `Unspecified` has no size and returns 0.
143pub const fn dtype_size(dtype: DType) -> usize {
144    match dtype {
145        DType::Unspecified => 0,
146        DType::Bool | DType::Uint8 | DType::Int8 => 1,
147        DType::Float16 | DType::Int16 | DType::Uint16 => 2,
148        DType::Int32 | DType::Uint32 | DType::Float32 => 4,
149        DType::Int64 | DType::Uint64 | DType::Float64 => 8,
150    }
151}
152
153#[cfg(test)]
154mod tests {
155    use super::*;
156
157    #[test]
158    fn test_dtype_i32_roundtrip() {
159        for dtype in DType::ALL {
160            let raw = i32::from(dtype);
161            assert_eq!(DType::try_from(raw), Ok(dtype));
162        }
163    }
164
165    #[test]
166    fn test_dtype_rejects_unknown_values() {
167        assert!(DType::try_from(-1).is_err());
168        for unused in [13, 20, 30, 99] {
169            assert!(
170                DType::try_from(unused).is_err(),
171                "{unused} should be invalid"
172            );
173        }
174    }
175
176    #[test]
177    fn test_dtype_name_roundtrip() {
178        for dtype in DType::ALL {
179            if dtype == DType::Unspecified {
180                continue;
181            }
182            assert_eq!(DType::from_name(dtype.name()), Some(dtype));
183            assert_eq!(dtype.to_string(), dtype.name());
184        }
185        assert_eq!(DType::Unspecified.name(), "unspecified");
186        assert_eq!(DType::from_name("unspecified"), None);
187        assert_eq!(DType::from_name("complex64"), None);
188        assert_eq!(DType::from_name("Float32"), None);
189    }
190
191    #[test]
192    fn test_dtype_size_table() {
193        let expected = [
194            (DType::Unspecified, 0),
195            (DType::Bool, 1),
196            (DType::Uint8, 1),
197            (DType::Int8, 1),
198            (DType::Float16, 2),
199            (DType::Int16, 2),
200            (DType::Uint16, 2),
201            (DType::Int32, 4),
202            (DType::Uint32, 4),
203            (DType::Float32, 4),
204            (DType::Int64, 8),
205            (DType::Uint64, 8),
206            (DType::Float64, 8),
207        ];
208        for (dtype, size) in expected {
209            assert_eq!(dtype_size(dtype), size, "size mismatch for {dtype:?}");
210        }
211    }
212}