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#[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 pub fn validate_with<V: ModelValidator + ?Sized>(
71 &self,
72 validator: &mut V,
73 ) -> Result<(), V::Error> {
74 validator.validate(self)
75 }
76
77 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 pub fn validate_for(&self, target: Target) -> Result<(), crate::SemanticError> {
99 crate::validate_semantics(self, target)
100 }
101
102 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
111pub 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#[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#[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#[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 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#[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 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#[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#[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 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#[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<'_> {}