1use core::fmt;
2
3use crate::generated::tosa as wire;
4
5#[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 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#[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#[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 pub const ALL: &'static [Self] = &[$(Self::$name),+];
92 };
93}
94
95#[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 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#[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}