1use core::fmt;
2
3use flatbuffers::Vector;
4
5use crate::generated::tosa as wire;
6use crate::{DType, Op};
7
8macro_rules! numeric_kind {
9 ($name:ident, $wire:ident, $($constant:ident),+ $(,)?) => {
10 #[derive(Clone, Copy, PartialEq, Eq, Hash)]
11 #[repr(transparent)]
12 pub struct $name(u32);
13
14 impl $name {
15 $(pub const $constant: Self = Self(wire::$wire::$constant.0);)+
16
17 pub const fn new(raw: u32) -> Self {
18 Self(raw)
19 }
20
21 pub const fn get(self) -> u32 {
22 self.0
23 }
24
25 pub fn name(self) -> Option<&'static str> {
26 wire::$wire(self.0).variant_name()
27 }
28 }
29
30 impl fmt::Debug for $name {
31 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
32 match self.name() {
33 Some(name) => formatter.write_str(name),
34 None => formatter.debug_tuple(stringify!($name)).field(&self.0).finish(),
35 }
36 }
37 }
38 };
39}
40
41numeric_kind!(
42 NanPropagationMode,
43 NanPropagationMode,
44 UNKNOWN,
45 PROPAGATE,
46 IGNORE,
47);
48numeric_kind!(ResizeMode, ResizeMode, UNKNOWN, NEAREST, BILINEAR);
49numeric_kind!(
50 RoundingMode,
51 RoundingMode,
52 UNKNOWN,
53 SINGLE_ROUND,
54 INEXACT_ROUND,
55 DOUBLE_ROUND,
56);
57
58#[derive(Clone, Copy)]
60pub struct I32List<'a>(Option<Vector<'a, i32>>);
61
62impl<'a> I32List<'a> {
63 pub fn len(self) -> usize {
64 match self.0 {
65 Some(values) => values.len(),
66 None => 0,
67 }
68 }
69
70 pub fn is_empty(self) -> bool {
71 self.len() == 0
72 }
73
74 pub fn get(self, index: usize) -> Option<i32> {
75 self.0
76 .and_then(|values| (index < values.len()).then(|| values.get(index)))
77 }
78
79 pub const fn iter(self) -> I32Values<'a> {
80 I32Values {
81 vector: self.0,
82 index: 0,
83 }
84 }
85}
86
87impl fmt::Debug for I32List<'_> {
88 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
89 formatter.debug_list().entries(self.iter()).finish()
90 }
91}
92
93#[derive(Clone)]
95pub struct I32Values<'a> {
96 vector: Option<Vector<'a, i32>>,
97 index: usize,
98}
99
100impl Iterator for I32Values<'_> {
101 type Item = i32;
102
103 fn next(&mut self) -> Option<Self::Item> {
104 let values = self.vector?;
105 if self.index >= values.len() {
106 return None;
107 }
108 let value = values.get(self.index);
109 self.index += 1;
110 Some(value)
111 }
112
113 fn size_hint(&self) -> (usize, Option<usize>) {
114 let remaining = self
115 .vector
116 .map_or(0, |values| values.len().saturating_sub(self.index));
117 (remaining, Some(remaining))
118 }
119}
120
121impl ExactSizeIterator for I32Values<'_> {}
122
123#[derive(Clone, Copy, Debug)]
128#[non_exhaustive]
129pub enum OpAttributes<'a> {
130 Empty {
131 op: Op,
132 },
133 ArgMax {
134 axis: i32,
135 nan_mode: NanPropagationMode,
136 },
137 AvgPool2d {
138 kernel: I32List<'a>,
139 stride: I32List<'a>,
140 pad: I32List<'a>,
141 acc_type: DType,
142 },
143 Conv2d {
144 pad: I32List<'a>,
145 stride: I32List<'a>,
146 dilation: I32List<'a>,
147 local_bound: bool,
148 acc_type: DType,
149 },
150 Conv3d {
151 pad: I32List<'a>,
152 stride: I32List<'a>,
153 dilation: I32List<'a>,
154 local_bound: bool,
155 acc_type: DType,
156 },
157 DepthwiseConv2d {
158 pad: I32List<'a>,
159 stride: I32List<'a>,
160 dilation: I32List<'a>,
161 local_bound: bool,
162 acc_type: DType,
163 },
164 Fft2d {
165 inverse: bool,
166 local_bound: bool,
167 },
168 MaxPool2d {
169 kernel: I32List<'a>,
170 stride: I32List<'a>,
171 pad: I32List<'a>,
172 nan_mode: NanPropagationMode,
173 },
174 Rfft2d {
175 local_bound: bool,
176 },
177 TransposeConv2d {
178 out_pad: I32List<'a>,
179 stride: I32List<'a>,
180 local_bound: bool,
181 acc_type: DType,
182 },
183 Clamp {
184 min_val: &'a [u8],
185 max_val: &'a [u8],
186 nan_mode: NanPropagationMode,
187 },
188 ArithmeticRightShift {
189 round: bool,
190 },
191 Maximum {
192 nan_mode: NanPropagationMode,
193 },
194 Minimum {
195 nan_mode: NanPropagationMode,
196 },
197 ReduceAll {
198 axis: i32,
199 },
200 ReduceAny {
201 axis: i32,
202 },
203 ReduceMax {
204 axis: i32,
205 nan_mode: NanPropagationMode,
206 },
207 ReduceMin {
208 axis: i32,
209 nan_mode: NanPropagationMode,
210 },
211 ReduceProduct {
212 axis: i32,
213 },
214 ReduceSum {
215 axis: i32,
216 },
217 Concat {
218 axis: i32,
219 },
220 Reverse {
221 axis: i32,
222 },
223 Transpose {
224 perms: I32List<'a>,
225 },
226 Resize {
227 mode: ResizeMode,
228 },
229 Rescale {
230 scale32: bool,
231 rounding_mode: RoundingMode,
232 per_channel: bool,
233 input_unsigned: bool,
234 output_unsigned: bool,
235 },
236 Custom {
237 operator_name: Option<&'a str>,
238 domain_name: Option<&'a str>,
239 implementation_attrs: &'a [u8],
240 },
241 CondIf {
242 then_graph: Option<&'a str>,
243 else_graph: Option<&'a str>,
244 },
245 WhileLoop {
246 cond_graph: Option<&'a str>,
247 body_graph: Option<&'a str>,
248 },
249}
250
251impl<'a> OpAttributes<'a> {
252 pub(crate) fn from_wire(operator: wire::TosaOperator<'a>) -> Self {
253 let op = Op::new(operator.op().0);
254 match op.get() {
255 1 => {
256 let value = operator
257 .attribute_as_arg_max_attribute()
258 .expect("validated ARGMAX attribute");
259 Self::ArgMax {
260 axis: value.axis(),
261 nan_mode: NanPropagationMode::new(value.nan_mode().0),
262 }
263 }
264 2 => {
265 let value = operator
266 .attribute_as_avg_pool_2d_attribute()
267 .expect("validated AVG_POOL2D attribute");
268 Self::AvgPool2d {
269 kernel: I32List(value.kernel()),
270 stride: I32List(value.stride()),
271 pad: I32List(value.pad()),
272 acc_type: DType::new(value.acc_type().0),
273 }
274 }
275 3 => {
276 let value = operator
277 .attribute_as_conv_2d_attribute()
278 .expect("validated CONV2D attribute");
279 Self::Conv2d {
280 pad: I32List(value.pad()),
281 stride: I32List(value.stride()),
282 dilation: I32List(value.dilation()),
283 local_bound: value.local_bound(),
284 acc_type: DType::new(value.acc_type().0),
285 }
286 }
287 4 => {
288 let value = operator
289 .attribute_as_conv_3d_attribute()
290 .expect("validated CONV3D attribute");
291 Self::Conv3d {
292 pad: I32List(value.pad()),
293 stride: I32List(value.stride()),
294 dilation: I32List(value.dilation()),
295 local_bound: value.local_bound(),
296 acc_type: DType::new(value.acc_type().0),
297 }
298 }
299 5 => {
300 let value = operator
301 .attribute_as_depthwise_conv_2d_attribute()
302 .expect("validated DEPTHWISE_CONV2D attribute");
303 Self::DepthwiseConv2d {
304 pad: I32List(value.pad()),
305 stride: I32List(value.stride()),
306 dilation: I32List(value.dilation()),
307 local_bound: value.local_bound(),
308 acc_type: DType::new(value.acc_type().0),
309 }
310 }
311 6 => {
312 let value = operator
313 .attribute_as_fft2d_attribute()
314 .expect("validated FFT2D attribute");
315 Self::Fft2d {
316 inverse: value.inverse(),
317 local_bound: value.local_bound(),
318 }
319 }
320 8 => {
321 let value = operator
322 .attribute_as_max_pool_2d_attribute()
323 .expect("validated MAX_POOL2D attribute");
324 Self::MaxPool2d {
325 kernel: I32List(value.kernel()),
326 stride: I32List(value.stride()),
327 pad: I32List(value.pad()),
328 nan_mode: NanPropagationMode::new(value.nan_mode().0),
329 }
330 }
331 9 => {
332 let value = operator
333 .attribute_as_rfft2d_attribute()
334 .expect("validated RFFT2D attribute");
335 Self::Rfft2d {
336 local_bound: value.local_bound(),
337 }
338 }
339 10 => {
340 let value = operator
341 .attribute_as_transpose_conv_2d_attribute()
342 .expect("validated TRANSPOSE_CONV2D attribute");
343 Self::TransposeConv2d {
344 out_pad: I32List(value.out_pad()),
345 stride: I32List(value.stride()),
346 local_bound: value.local_bound(),
347 acc_type: DType::new(value.acc_type().0),
348 }
349 }
350 11 => {
351 let value = operator
352 .attribute_as_clamp_attribute()
353 .expect("validated CLAMP attribute");
354 Self::Clamp {
355 min_val: value.min_val().map_or(&[], |bytes| bytes.bytes()),
356 max_val: value.max_val().map_or(&[], |bytes| bytes.bytes()),
357 nan_mode: NanPropagationMode::new(value.nan_mode().0),
358 }
359 }
360 16 => {
361 let value = operator
362 .attribute_as_arithmetic_right_shift_attribute()
363 .expect("validated ARITHMETIC_RIGHT_SHIFT attribute");
364 Self::ArithmeticRightShift {
365 round: value.round(),
366 }
367 }
368 26 => {
369 let value = operator
370 .attribute_as_maximum_attribute()
371 .expect("validated MAXIMUM attribute");
372 Self::Maximum {
373 nan_mode: NanPropagationMode::new(value.nan_mode().0),
374 }
375 }
376 27 => {
377 let value = operator
378 .attribute_as_minimum_attribute()
379 .expect("validated MINIMUM attribute");
380 Self::Minimum {
381 nan_mode: NanPropagationMode::new(value.nan_mode().0),
382 }
383 }
384 49 => {
385 let value = operator
386 .attribute_as_reduce_all_attribute()
387 .expect("validated REDUCE_ALL attribute");
388 Self::ReduceAll { axis: value.axis() }
389 }
390 50 => {
391 let value = operator
392 .attribute_as_reduce_any_attribute()
393 .expect("validated REDUCE_ANY attribute");
394 Self::ReduceAny { axis: value.axis() }
395 }
396 51 => {
397 let value = operator
398 .attribute_as_reduce_max_attribute()
399 .expect("validated REDUCE_MAX attribute");
400 Self::ReduceMax {
401 axis: value.axis(),
402 nan_mode: NanPropagationMode::new(value.nan_mode().0),
403 }
404 }
405 52 => {
406 let value = operator
407 .attribute_as_reduce_min_attribute()
408 .expect("validated REDUCE_MIN attribute");
409 Self::ReduceMin {
410 axis: value.axis(),
411 nan_mode: NanPropagationMode::new(value.nan_mode().0),
412 }
413 }
414 53 => {
415 let value = operator
416 .attribute_as_reduce_product_attribute()
417 .expect("validated REDUCE_PRODUCT attribute");
418 Self::ReduceProduct { axis: value.axis() }
419 }
420 54 => {
421 let value = operator
422 .attribute_as_reduce_sum_attribute()
423 .expect("validated REDUCE_SUM attribute");
424 Self::ReduceSum { axis: value.axis() }
425 }
426 55 => {
427 let value = operator
428 .attribute_as_concat_attribute()
429 .expect("validated CONCAT attribute");
430 Self::Concat { axis: value.axis() }
431 }
432 58 => {
433 let value = operator
434 .attribute_as_reverse_attribute()
435 .expect("validated REVERSE attribute");
436 Self::Reverse { axis: value.axis() }
437 }
438 61 => {
439 let value = operator
440 .attribute_as_transpose_attribute()
441 .expect("validated TRANSPOSE attribute");
442 Self::Transpose {
443 perms: I32List(value.perms()),
444 }
445 }
446 64 => {
447 let value = operator
448 .attribute_as_resize_attribute()
449 .expect("validated RESIZE attribute");
450 Self::Resize {
451 mode: ResizeMode::new(value.mode().0),
452 }
453 }
454 66 => {
455 let value = operator
456 .attribute_as_rescale_attribute()
457 .expect("validated RESCALE attribute");
458 Self::Rescale {
459 scale32: value.scale32(),
460 rounding_mode: RoundingMode::new(value.rounding_mode().0),
461 per_channel: value.per_channel(),
462 input_unsigned: value.input_unsigned(),
463 output_unsigned: value.output_unsigned(),
464 }
465 }
466 69 => {
467 let value = operator
468 .attribute_as_custom_attribute()
469 .expect("validated CUSTOM attribute");
470 Self::Custom {
471 operator_name: value.operator_name(),
472 domain_name: value.domain_name(),
473 implementation_attrs: value
474 .implementation_attrs()
475 .map_or(&[], |bytes| bytes.bytes()),
476 }
477 }
478 70 => {
479 let value = operator
480 .attribute_as_cond_if_attribute()
481 .expect("validated COND_IF attribute");
482 Self::CondIf {
483 then_graph: value.then_graph(),
484 else_graph: value.else_graph(),
485 }
486 }
487 71 => {
488 let value = operator
489 .attribute_as_while_loop_attribute()
490 .expect("validated WHILE_LOOP attribute");
491 Self::WhileLoop {
492 cond_graph: value.cond_graph(),
493 body_graph: value.body_graph(),
494 }
495 }
496 _ => Self::Empty { op },
497 }
498 }
499}