Skip to main content

virtio_accel_tosa/
view.rs

1use crate::artifact::{ArtifactError, Target};
2use crate::generated::tosa as wire;
3use crate::{AttributeKind, DType, Op, OpAttributes, Stats, Version};
4use flatbuffers::{ForwardsUOffset, Vector};
5use virtio_accel_core::{ArtifactRef, BackendError, ByteSource};
6
7type WireRegions<'a> = Vector<'a, ForwardsUOffset<wire::TosaRegion<'a>>>;
8type WireBlocks<'a> = Vector<'a, ForwardsUOffset<wire::TosaBasicBlock<'a>>>;
9type WireTensors<'a> = Vector<'a, ForwardsUOffset<wire::TosaTensor<'a>>>;
10type WireShapes<'a> = Vector<'a, ForwardsUOffset<wire::TosaShape<'a>>>;
11type WireOperators<'a> = Vector<'a, ForwardsUOffset<wire::TosaOperator<'a>>>;
12type WireStrings<'a> = Vector<'a, ForwardsUOffset<&'a str>>;
13
14/// A verified, zero-copy TOSA graph.
15#[derive(Clone, Copy)]
16pub struct Model<'a> {
17    pub(crate) graph: wire::TosaGraph<'a>,
18    pub(crate) bytes: &'a [u8],
19    pub(crate) version: Version,
20    pub(crate) stats: Stats,
21    pub(crate) source: SliceSource<'a>,
22}
23
24#[derive(Clone, Copy, Debug)]
25pub(crate) struct SliceSource<'a>(pub(crate) &'a [u8]);
26
27impl ByteSource for SliceSource<'_> {
28    fn len(&self) -> u64 {
29        self.0.len() as u64
30    }
31
32    fn read_at(&self, offset: u64, target: &mut [u8]) -> Result<(), BackendError> {
33        ByteSource::read_at(self.0, offset, target)
34    }
35
36    fn as_contiguous(&self) -> Option<&[u8]> {
37        Some(self.0)
38    }
39}
40
41impl core::fmt::Debug for Model<'_> {
42    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
43        formatter
44            .debug_struct("Model")
45            .field("version", &self.version)
46            .field("stats", &self.stats)
47            .field("bytes", &self.bytes.len())
48            .finish()
49    }
50}
51
52impl<'a> Model<'a> {
53    pub const fn version(&self) -> Version {
54        self.version
55    }
56
57    pub const fn stats(&self) -> Stats {
58        self.stats
59    }
60
61    pub const fn as_bytes(&self) -> &'a [u8] {
62        self.bytes
63    }
64
65    pub fn regions(&self) -> Regions<'a> {
66        Regions::new(self.graph.regions())
67    }
68
69    /// Run an operator- or provider-specific validation pass over this safe graph view.
70    pub fn validate_with<V: ModelValidator + ?Sized>(
71        &self,
72        validator: &mut V,
73    ) -> Result<(), V::Error> {
74        validator.validate(self)
75    }
76
77    /// Wrap these exact validated bytes in the standard `virtio-accel` artifact envelope.
78    pub fn artifact_ref(
79        &self,
80        target: Target,
81        resident_bytes: u64,
82    ) -> Result<ArtifactRef<'_>, ArtifactError> {
83        if target.version != self.version {
84            return Err(ArtifactError::VersionMismatch {
85                model: self.version,
86                target: target.version,
87            });
88        }
89        Ok(ArtifactRef {
90            format: crate::ARTIFACT_FORMAT,
91            target: target.to_identity(),
92            payload: &self.source,
93            resident_bytes,
94        })
95    }
96
97    /// Apply the complete stable TOSA semantic pass for a device-neutral target.
98    pub fn validate_for(&self, target: Target) -> Result<(), crate::SemanticError> {
99        crate::validate_semantics(self, target)
100    }
101
102    /// Validate once and build a compact provider-neutral lowering plan over this borrowed model.
103    pub fn analyze_for(
104        &self,
105        target: Target,
106    ) -> Result<crate::TosaAnalysis<'a>, crate::AnalysisError> {
107        crate::TosaAnalysis::build(self, target)
108    }
109}
110
111/// Extension point for semantic, profile, extension, or backend-capability validators.
112pub trait ModelValidator {
113    type Error;
114
115    fn validate(&mut self, model: &Model<'_>) -> Result<(), Self::Error>;
116}
117
118macro_rules! table_iter {
119    ($iterator:ident, $item:ident, $wire:ty, $wrapper:expr) => {
120        #[derive(Clone)]
121        pub struct $iterator<'a> {
122            vector: Option<$wire>,
123            index: usize,
124        }
125
126        impl<'a> $iterator<'a> {
127            fn new(vector: Option<$wire>) -> Self {
128                Self { vector, index: 0 }
129            }
130        }
131
132        impl<'a> Iterator for $iterator<'a> {
133            type Item = $item<'a>;
134
135            fn next(&mut self) -> Option<Self::Item> {
136                let vector = self.vector?;
137                if self.index >= vector.len() {
138                    return None;
139                }
140                let item = vector.get(self.index);
141                self.index += 1;
142                Some($wrapper(item))
143            }
144
145            fn size_hint(&self) -> (usize, Option<usize>) {
146                let remaining = self
147                    .vector
148                    .map_or(0, |vector| vector.len().saturating_sub(self.index));
149                (remaining, Some(remaining))
150            }
151        }
152
153        impl ExactSizeIterator for $iterator<'_> {}
154    };
155}
156
157/// Borrowed region view.
158#[derive(Clone, Copy, Debug)]
159pub struct Region<'a>(wire::TosaRegion<'a>);
160
161impl<'a> Region<'a> {
162    pub fn name(&self) -> &'a str {
163        self.0.name().expect("validated region name")
164    }
165
166    pub fn blocks(&self) -> BasicBlocks<'a> {
167        BasicBlocks::new(self.0.blocks())
168    }
169}
170
171table_iter!(Regions, Region, WireRegions<'a>, Region);
172
173/// Borrowed basic-block view.
174#[derive(Clone, Copy, Debug)]
175pub struct BasicBlock<'a>(wire::TosaBasicBlock<'a>);
176
177impl<'a> BasicBlock<'a> {
178    pub fn name(&self) -> &'a str {
179        self.0.name().expect("validated block name")
180    }
181
182    pub fn tensors(&self) -> Tensors<'a> {
183        Tensors::new(self.0.tensors())
184    }
185
186    pub fn shapes(&self) -> Shapes<'a> {
187        Shapes::new(self.0.shapes())
188    }
189
190    pub fn operators(&self) -> Operators<'a> {
191        Operators::new(self.0.operators())
192    }
193
194    pub fn inputs(&self) -> StringList<'a> {
195        StringList::new(self.0.inputs())
196    }
197
198    pub fn outputs(&self) -> StringList<'a> {
199        StringList::new(self.0.outputs())
200    }
201}
202
203table_iter!(BasicBlocks, BasicBlock, WireBlocks<'a>, BasicBlock);
204
205/// Borrowed tensor view.
206#[derive(Clone, Copy, Debug)]
207pub struct Tensor<'a>(wire::TosaTensor<'a>);
208
209impl<'a> Tensor<'a> {
210    pub fn name(&self) -> &'a str {
211        self.0.name().expect("validated tensor name")
212    }
213
214    pub fn dtype(&self) -> DType {
215        DType::new(self.0.type_().0)
216    }
217
218    pub fn rank(&self) -> Option<usize> {
219        if self.0.is_unranked() {
220            None
221        } else {
222            Some(self.0.shape().map_or(0, |shape| shape.len()))
223        }
224    }
225
226    pub fn dimensions(&self) -> impl Iterator<Item = i32> + 'a {
227        self.0.shape().into_iter().flat_map(|shape| shape.iter())
228    }
229
230    /// One ranked dimension without constructing or advancing an iterator.
231    pub fn dimension(&self, index: usize) -> Option<i32> {
232        self.0
233            .shape()
234            .filter(|_| !self.0.is_unranked())
235            .and_then(|shape| (index < shape.len()).then(|| shape.get(index)))
236    }
237
238    pub fn data(&self) -> &'a [u8] {
239        self.0.data().map_or(&[], |data| data.bytes())
240    }
241
242    pub fn is_variable(&self) -> bool {
243        self.0.variable()
244    }
245
246    pub fn variable_name(&self) -> Option<&'a str> {
247        self.0.variable_name()
248    }
249
250    pub fn external_data_range(&self) -> Option<(u64, u64)> {
251        let size = self.0.size();
252        (size != 0).then_some((self.0.offset(), size))
253    }
254}
255
256table_iter!(Tensors, Tensor, WireTensors<'a>, Tensor);
257
258/// Borrowed shape-value view.
259#[derive(Clone, Copy, Debug)]
260pub struct Shape<'a>(wire::TosaShape<'a>);
261
262impl<'a> Shape<'a> {
263    pub fn name(&self) -> &'a str {
264        self.0.name().expect("validated shape name")
265    }
266
267    pub fn rank(&self) -> u32 {
268        self.0.rank()
269    }
270
271    pub fn data(&self) -> &'a [u8] {
272        self.0.data().map_or(&[], |data| data.bytes())
273    }
274
275    /// Decoded shape values, or `None` for a nonempty intermediate shape without constant data.
276    /// A rank-zero shape is the constant empty list and therefore returns an empty iterator.
277    pub fn values(&self) -> Option<ShapeValues<'a>> {
278        let data = self.data();
279        (!data.is_empty() || self.rank() == 0).then_some(ShapeValues {
280            data,
281            index: 0,
282            len: self.rank() as usize,
283        })
284    }
285}
286
287table_iter!(Shapes, Shape, WireShapes<'a>, Shape);
288
289/// Exact-size iterator over little-endian `i64` shape values.
290#[derive(Clone)]
291pub struct ShapeValues<'a> {
292    data: &'a [u8],
293    index: usize,
294    len: usize,
295}
296
297impl Iterator for ShapeValues<'_> {
298    type Item = i64;
299
300    fn next(&mut self) -> Option<Self::Item> {
301        if self.index >= self.len {
302            return None;
303        }
304        let start = self.index * core::mem::size_of::<i64>();
305        let bytes: [u8; 8] = self.data[start..start + 8]
306            .try_into()
307            .expect("validated shape data");
308        self.index += 1;
309        Some(i64::from_le_bytes(bytes))
310    }
311
312    fn size_hint(&self) -> (usize, Option<usize>) {
313        let remaining = self.len.saturating_sub(self.index);
314        (remaining, Some(remaining))
315    }
316}
317
318impl ExactSizeIterator for ShapeValues<'_> {}
319
320/// Borrowed operator view.
321#[derive(Clone, Copy, Debug)]
322pub struct Operator<'a>(wire::TosaOperator<'a>);
323
324impl<'a> Operator<'a> {
325    pub fn op(&self) -> Op {
326        Op::new(self.0.op().0)
327    }
328
329    pub fn attribute_kind(&self) -> AttributeKind {
330        AttributeKind::new(self.0.attribute_type().0)
331    }
332
333    /// Safe view of every field in this operator's stable TOSA 1.0 attribute table.
334    pub fn attributes(&self) -> OpAttributes<'a> {
335        OpAttributes::from_wire(self.0)
336    }
337
338    pub fn inputs(&self) -> StringList<'a> {
339        StringList::new(self.0.inputs())
340    }
341
342    pub fn outputs(&self) -> StringList<'a> {
343        StringList::new(self.0.outputs())
344    }
345
346    pub fn location(&self) -> Option<&'a str> {
347        self.0.location().and_then(|location| location.text())
348    }
349}
350
351table_iter!(Operators, Operator, WireOperators<'a>, Operator);
352
353/// Iterator over borrowed TOSA symbol names.
354#[derive(Clone)]
355pub struct StringList<'a> {
356    vector: Option<WireStrings<'a>>,
357    index: usize,
358}
359
360impl<'a> StringList<'a> {
361    fn new(vector: Option<WireStrings<'a>>) -> Self {
362        Self { vector, index: 0 }
363    }
364}
365
366impl<'a> Iterator for StringList<'a> {
367    type Item = &'a str;
368
369    fn next(&mut self) -> Option<Self::Item> {
370        let vector = self.vector?;
371        if self.index >= vector.len() {
372            return None;
373        }
374        let item = vector.get(self.index);
375        self.index += 1;
376        Some(item)
377    }
378
379    fn size_hint(&self) -> (usize, Option<usize>) {
380        let remaining = self
381            .vector
382            .map_or(0, |vector| vector.len().saturating_sub(self.index));
383        (remaining, Some(remaining))
384    }
385}
386
387impl ExactSizeIterator for StringList<'_> {}