Skip to main content

virtio_accel_tosa/
types.rs

1use core::fmt;
2
3use crate::generated::tosa as wire;
4
5/// Stable TOSA graph version.
6#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
7pub struct Version {
8    pub major: u16,
9    pub minor: u16,
10    pub patch: u16,
11}
12
13impl Version {
14    /// Version accepted by the default validator.
15    pub const TOSA_1_0: Self = Self {
16        major: 1,
17        minor: 0,
18        patch: 0,
19    };
20
21    pub const fn new(major: u16, minor: u16, patch: u16) -> Self {
22        Self {
23            major,
24            minor,
25            patch,
26        }
27    }
28}
29
30/// A TOSA tensor data-type discriminant.
31///
32/// Raw values remain representable so policy and diagnostic utilities can discuss newer schemas.
33/// Successfully parsed models contain only values accepted by [`DType::is_tosa_1_0`].
34#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
35#[repr(transparent)]
36pub struct DType(u32);
37
38impl DType {
39    pub const BOOL: Self = Self(1);
40    pub const INT4: Self = Self(2);
41    pub const INT8: Self = Self(3);
42    pub const INT16: Self = Self(4);
43    pub const INT32: Self = Self(5);
44    pub const INT48: Self = Self(6);
45    pub const FP32: Self = Self(7);
46    pub const FP16: Self = Self(8);
47    pub const BF16: Self = Self(9);
48    pub const SHAPE: Self = Self(10);
49    pub const FP8E4M3: Self = Self(11);
50    pub const FP8E5M2: Self = Self(12);
51
52    pub const fn new(raw: u32) -> Self {
53        Self(raw)
54    }
55
56    pub const fn get(self) -> u32 {
57        self.0
58    }
59
60    pub const fn is_tosa_1_0(self) -> bool {
61        self.0 >= Self::BOOL.0 && self.0 <= Self::FP8E5M2.0
62    }
63
64    pub fn name(self) -> Option<&'static str> {
65        wire::DType(self.0).variant_name()
66    }
67}
68
69impl fmt::Debug for DType {
70    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
71        match self.name() {
72            Some(name) => formatter.write_str(name),
73            None => formatter.debug_tuple("DType").field(&self.0).finish(),
74        }
75    }
76}
77
78/// A TOSA operator discriminant.
79///
80/// Constants cover the stable TOSA 1.0 operator set. Raw values above [`Op::CONST_SHAPE`] belong
81/// to schema additions newer than TOSA 1.0 and are rejected by [`crate::parse`].
82#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
83#[repr(transparent)]
84pub struct Op(u32);
85
86macro_rules! op_constants {
87    ($($name:ident = $value:literal),+ $(,)?) => {
88        $(pub const $name: Self = Self($value);)+
89
90        /// Every operator in the stable TOSA 1.0 serialization set, in wire order.
91        pub const ALL: &'static [Self] = &[$(Self::$name),+];
92    };
93}
94
95/// Accepted serialized operand counts for an operator.
96#[derive(Clone, Copy, Debug, PartialEq, Eq)]
97pub struct Arity {
98    pub min_inputs: usize,
99    pub max_inputs: Option<usize>,
100    pub min_outputs: usize,
101    pub max_outputs: Option<usize>,
102    pub matching_input_output_counts: bool,
103}
104
105impl Arity {
106    pub const fn exact(inputs: usize, outputs: usize) -> Self {
107        Self {
108            min_inputs: inputs,
109            max_inputs: Some(inputs),
110            min_outputs: outputs,
111            max_outputs: Some(outputs),
112            matching_input_output_counts: false,
113        }
114    }
115
116    pub const fn accepts(self, inputs: usize, outputs: usize) -> bool {
117        inputs >= self.min_inputs
118            && match self.max_inputs {
119                Some(maximum) => inputs <= maximum,
120                None => true,
121            }
122            && outputs >= self.min_outputs
123            && match self.max_outputs {
124                Some(maximum) => outputs <= maximum,
125                None => true,
126            }
127            && (!self.matching_input_output_counts || inputs == outputs)
128    }
129}
130
131impl Op {
132    op_constants! {
133        ARGMAX = 1,
134        AVG_POOL2D = 2,
135        CONV2D = 3,
136        CONV3D = 4,
137        DEPTHWISE_CONV2D = 5,
138        FFT2D = 6,
139        MATMUL = 7,
140        MAX_POOL2D = 8,
141        RFFT2D = 9,
142        TRANSPOSE_CONV2D = 10,
143        CLAMP = 11,
144        ERF = 12,
145        SIGMOID = 13,
146        TANH = 14,
147        ADD = 15,
148        ARITHMETIC_RIGHT_SHIFT = 16,
149        BITWISE_AND = 17,
150        BITWISE_OR = 18,
151        BITWISE_XOR = 19,
152        INTDIV = 20,
153        LOGICAL_AND = 21,
154        LOGICAL_LEFT_SHIFT = 22,
155        LOGICAL_RIGHT_SHIFT = 23,
156        LOGICAL_OR = 24,
157        LOGICAL_XOR = 25,
158        MAXIMUM = 26,
159        MINIMUM = 27,
160        MUL = 28,
161        POW = 29,
162        SUB = 30,
163        TABLE = 31,
164        ABS = 32,
165        BITWISE_NOT = 33,
166        CEIL = 34,
167        CLZ = 35,
168        COS = 36,
169        EXP = 37,
170        FLOOR = 38,
171        LOG = 39,
172        LOGICAL_NOT = 40,
173        NEGATE = 41,
174        RECIPROCAL = 42,
175        RSQRT = 43,
176        SIN = 44,
177        SELECT = 45,
178        EQUAL = 46,
179        GREATER = 47,
180        GREATER_EQUAL = 48,
181        REDUCE_ALL = 49,
182        REDUCE_ANY = 50,
183        REDUCE_MAX = 51,
184        REDUCE_MIN = 52,
185        REDUCE_PRODUCT = 53,
186        REDUCE_SUM = 54,
187        CONCAT = 55,
188        PAD = 56,
189        RESHAPE = 57,
190        REVERSE = 58,
191        SLICE = 59,
192        TILE = 60,
193        TRANSPOSE = 61,
194        GATHER = 62,
195        SCATTER = 63,
196        RESIZE = 64,
197        CAST = 65,
198        RESCALE = 66,
199        CONST = 67,
200        IDENTITY = 68,
201        CUSTOM = 69,
202        COND_IF = 70,
203        WHILE_LOOP = 71,
204        VARIABLE = 72,
205        VARIABLE_WRITE = 73,
206        VARIABLE_READ = 74,
207        CONST_SHAPE = 75,
208    }
209
210    pub const fn new(raw: u32) -> Self {
211        Self(raw)
212    }
213
214    pub const fn get(self) -> u32 {
215        self.0
216    }
217
218    pub const fn is_tosa_1_0(self) -> bool {
219        self.0 >= Self::ARGMAX.0 && self.0 <= Self::CONST_SHAPE.0
220    }
221
222    pub fn name(self) -> Option<&'static str> {
223        wire::Op(self.0).variant_name()
224    }
225
226    /// Stable TOSA 1.0 serialized operand-count contract.
227    pub const fn arity(self) -> Option<Arity> {
228        let exact = match self.0 {
229            1 => (1, 1),
230            2 => (3, 1),
231            3..=5 => (5, 1),
232            6 => (2, 2),
233            7 => (4, 1),
234            8 => (1, 1),
235            9 => (1, 2),
236            10 => (5, 1),
237            11..=14 => (1, 1),
238            15..=27 => (2, 1),
239            28 => (3, 1),
240            29..=31 => (2, 1),
241            32..=40 => (1, 1),
242            41 => (3, 1),
243            42..=44 => (1, 1),
244            45 => (3, 1),
245            46..=48 => (2, 1),
246            49..=54 => (1, 1),
247            56 => (3, 1),
248            57 => (2, 1),
249            58 => (1, 1),
250            59 => (3, 1),
251            60 => (2, 1),
252            61 => (1, 1),
253            62 => (2, 1),
254            63 => (3, 1),
255            64 => (4, 1),
256            65 => (1, 1),
257            66 => (5, 1),
258            67 => (0, 1),
259            68 => (1, 1),
260            72 => (0, 0),
261            73 => (1, 0),
262            74 => (0, 1),
263            75 => (0, 1),
264            55 | 69..=71 => {
265                return Some(match self.0 {
266                    55 => Arity {
267                        min_inputs: 1,
268                        max_inputs: None,
269                        min_outputs: 1,
270                        max_outputs: Some(1),
271                        matching_input_output_counts: false,
272                    },
273                    69 => Arity {
274                        min_inputs: 0,
275                        max_inputs: None,
276                        min_outputs: 0,
277                        max_outputs: None,
278                        matching_input_output_counts: false,
279                    },
280                    70 => Arity {
281                        min_inputs: 1,
282                        max_inputs: None,
283                        min_outputs: 0,
284                        max_outputs: None,
285                        matching_input_output_counts: false,
286                    },
287                    71 => Arity {
288                        min_inputs: 0,
289                        max_inputs: None,
290                        min_outputs: 0,
291                        max_outputs: None,
292                        matching_input_output_counts: true,
293                    },
294                    _ => unreachable!(),
295                });
296            }
297            _ => return None,
298        };
299        Some(Arity::exact(exact.0, exact.1))
300    }
301}
302
303impl fmt::Debug for Op {
304    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
305        match self.name() {
306            Some(name) => formatter.write_str(name),
307            None => formatter.debug_tuple("Op").field(&self.0).finish(),
308        }
309    }
310}
311
312/// Discriminant of an operator's serialized attribute table.
313#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
314#[repr(transparent)]
315pub struct AttributeKind(u8);
316
317impl AttributeKind {
318    pub const fn new(raw: u8) -> Self {
319        Self(raw)
320    }
321
322    pub const fn get(self) -> u8 {
323        self.0
324    }
325
326    pub fn name(self) -> Option<&'static str> {
327        wire::Attribute(self.0).variant_name()
328    }
329}
330
331impl fmt::Debug for AttributeKind {
332    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
333        match self.name() {
334            Some(name) => formatter.write_str(name),
335            None => formatter
336                .debug_tuple("AttributeKind")
337                .field(&self.0)
338                .finish(),
339        }
340    }
341}