Skip to main content

virtio_accel_tosa/
attribute.rs

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/// Borrowed FlatBuffers vector of little-endian `i32` attribute values.
59#[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/// Exact-size iterator returned by [`I32List::iter`].
94#[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/// Safe, exhaustive view of the attribute payload used by a stable TOSA 1.0 operator.
124///
125/// Operators whose schema table has no fields return [`Self::Empty`]; [`Self::Empty::op`] still
126/// identifies the exact table. Vector fields remain borrowed and are decoded without allocation.
127#[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}