Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/onnx.rs

13.9 KiB, 1 run

created by r1870400018:19744, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1//! A reader for the subset of the ONNX wire format a small convolutional
2//! network uses.
3//!
4//! ONNX is protocol buffers, and protocol buffers can be walked without a
5//! schema: every field carries its number and its wire type. This module reads
6//! only the fields it needs -- nodes, their attributes, and the initialisers
7//! that hold the weights -- and ignores the rest, so it is a few hundred lines
8//! rather than a generated library.
9//!
10//! Nothing here interprets the graph. [`crate::graph`] does that.
11
12use oxedyne_fe2o3_core::prelude::*;
13
14/// Protocol buffer wire types this reader understands.
15const WIRE_VARINT: u8 = 0;
16const WIRE_I64: u8 = 1;
17const WIRE_LEN: u8 = 2;
18const WIRE_I32: u8 = 5;
19
20/// ONNX tensor element types this reader understands.
21const DT_FLOAT: i64 = 1;
22const DT_INT64: i64 = 7;
23
24/// A cursor over a protocol buffer message.
25struct Reader<'a> {
26 /// The bytes of one message.
27 buf: &'a [u8],
28 /// Read position within `buf`.
29 pos: usize,
30}
31
32/// One field of a protocol buffer message, as read off the wire.
33enum Field<'a> {
34 /// A base-128 integer.
35 Varint(u64),
36 /// A length-delimited payload: a string, a submessage or packed values.
37 Bytes(&'a [u8]),
38 /// A fixed thirty-two bit value.
39 Fixed32([u8; 4]),
40 /// A fixed sixty-four bit value, which no field this reader wants uses.
41 Fixed64,
42}
43
44impl<'a> Reader<'a> {
45 /// Starts a cursor over a message.
46 fn new(buf: &'a [u8]) -> Self {
47 Self { buf, pos: 0 }
48 }
49
50 /// Whether every byte has been consumed.
51 fn done(&self) -> bool {
52 self.pos >= self.buf.len()
53 }
54
55 /// Reads one base-128 integer.
56 fn varint(&mut self) -> Outcome<u64> {
57 let mut r = 0u64;
58 let mut shift = 0u32;
59 loop {
60 if self.pos >= self.buf.len() {
61 return Err(err!("A base-128 integer runs past the end of the message.";
62 Invalid, Input, Decode));
63 }
64 let b = self.buf[self.pos];
65 self.pos += 1;
66 if shift >= 64 {
67 return Err(err!("A base-128 integer is wider than sixty-four bits.";
68 Invalid, Input, Decode));
69 }
70 r |= ((b & 0x7f) as u64) << shift;
71 if b & 0x80 == 0 {
72 return Ok(r);
73 }
74 shift += 7;
75 }
76 }
77
78 /// Reads the next field, answering its number and payload.
79 fn next(&mut self) -> Outcome<(u64, Field<'a>)> {
80 let key = res!(self.varint());
81 let num = key >> 3;
82 let wire = (key & 7) as u8;
83 let f = match wire {
84 WIRE_VARINT => Field::Varint(res!(self.varint())),
85 WIRE_LEN => {
86 let n = res!(self.varint()) as usize;
87 let end = match self.pos.checked_add(n) {
88 Some(e) if e <= self.buf.len() => e,
89 _ => return Err(err!(
90 "A length-delimited field of {} bytes at {} runs past the end of a \
91 message of {} bytes.", n, self.pos, self.buf.len();
92 Invalid, Input, Decode)),
93 };
94 let s = &self.buf[self.pos..end];
95 self.pos = end;
96 Field::Bytes(s)
97 },
98 WIRE_I32 => {
99 if self.pos + 4 > self.buf.len() {
100 return Err(err!("A thirty-two bit field runs past the end of the message.";
101 Invalid, Input, Decode));
102 }
103 let mut a = [0u8; 4];
104 a.copy_from_slice(&self.buf[self.pos..self.pos + 4]);
105 self.pos += 4;
106 Field::Fixed32(a)
107 },
108 WIRE_I64 => {
109 if self.pos + 8 > self.buf.len() {
110 return Err(err!("A sixty-four bit field runs past the end of the message.";
111 Invalid, Input, Decode));
112 }
113 self.pos += 8;
114 Field::Fixed64
115 },
116 other => return Err(err!(
117 "Wire type {} is not one this reader knows.", other; Invalid, Input, Decode)),
118 };
119 Ok((num, f))
120 }
121}
122
123/// Reads a payload as a UTF-8 string.
124fn as_str(b: &[u8]) -> Outcome<String> {
125 match core::str::from_utf8(b) {
126 Ok(s) => Ok(s.to_string()),
127 Err(e) => Err(err!(e, "A name in the model is not valid UTF-8."; Invalid, Input, Decode)),
128 }
129}
130
131/// Reads a payload as packed base-128 integers.
132fn packed_varints(b: &[u8]) -> Outcome<Vec<i64>> {
133 let mut r = Reader::new(b);
134 let mut out = Vec::new();
135 while !r.done() {
136 out.push(res!(r.varint()) as i64);
137 }
138 Ok(out)
139}
140
141/// Reads a payload as little-endian `f32` values.
142fn le_f32(b: &[u8]) -> Outcome<Vec<f32>> {
143 if b.len() % 4 != 0 {
144 return Err(err!(
145 "A block of {} bytes does not divide into four-byte floats.", b.len();
146 Invalid, Input, Decode));
147 }
148 let mut out = Vec::with_capacity(b.len() / 4);
149 for c in b.chunks_exact(4) {
150 out.push(f32::from_le_bytes([c[0], c[1], c[2], c[3]]));
151 }
152 Ok(out)
153}
154
155/// Reads a payload as little-endian `i64` values.
156fn le_i64(b: &[u8]) -> Outcome<Vec<i64>> {
157 if b.len() % 8 != 0 {
158 return Err(err!(
159 "A block of {} bytes does not divide into eight-byte integers.", b.len();
160 Invalid, Input, Decode));
161 }
162 let mut out = Vec::with_capacity(b.len() / 8);
163 for c in b.chunks_exact(8) {
164 let mut a = [0u8; 8];
165 a.copy_from_slice(c);
166 out.push(i64::from_le_bytes(a));
167 }
168 Ok(out)
169}
170
171/// The value of one node attribute.
172#[derive(Clone, Debug)]
173pub enum Attr {
174 /// A single integer.
175 Int(i64),
176 /// A single float.
177 Float(f32),
178 /// A single string.
179 Str(String),
180 /// A list of integers.
181 Ints(Vec<i64>),
182 /// A list of floats.
183 Floats(Vec<f32>),
184}
185
186impl Attr {
187 /// Reads the attribute as a single integer.
188 pub fn int(&self) -> Outcome<i64> {
189 match self {
190 Self::Int(v) => Ok(*v),
191 other => Err(err!("An integer attribute was expected, found {:?}.", other;
192 Invalid, Input, Mismatch)),
193 }
194 }
195
196 /// Reads the attribute as a single float.
197 pub fn float(&self) -> Outcome<f32> {
198 match self {
199 Self::Float(v) => Ok(*v),
200 other => Err(err!("A float attribute was expected, found {:?}.", other;
201 Invalid, Input, Mismatch)),
202 }
203 }
204
205 /// Reads the attribute as a list of integers.
206 pub fn ints(&self) -> Outcome<&[i64]> {
207 match self {
208 Self::Ints(v) => Ok(v),
209 other => Err(err!("A list of integers was expected, found {:?}.", other;
210 Invalid, Input, Mismatch)),
211 }
212 }
213
214 /// Reads the attribute as a string.
215 pub fn text(&self) -> Outcome<&str> {
216 match self {
217 Self::Str(v) => Ok(v),
218 other => Err(err!("A string attribute was expected, found {:?}.", other;
219 Invalid, Input, Mismatch)),
220 }
221 }
222}
223
224/// An initialiser, in whichever element type the model stored it.
225#[derive(Clone, Debug)]
226pub enum Init {
227 /// Thirty-two bit floats -- weights, biases, scales.
228 F32 {
229 /// Extent along each axis.
230 dims: Vec<usize>,
231 /// Values, row-major.
232 data: Vec<f32>,
233 },
234 /// Sixty-four bit integers -- shapes, axes, permutations.
235 I64 {
236 /// Extent along each axis.
237 dims: Vec<usize>,
238 /// Values, row-major.
239 data: Vec<i64>,
240 },
241}
242
243impl Init {
244 /// Extent along each axis.
245 pub fn dims(&self) -> &[usize] {
246 match self {
247 Self::F32 { dims, .. } => dims,
248 Self::I64 { dims, .. } => dims,
249 }
250 }
251
252 /// Reads the initialiser as floats.
253 pub fn floats(&self) -> Outcome<&[f32]> {
254 match self {
255 Self::F32 { data, .. } => Ok(data),
256 Self::I64 { dims, .. } => Err(err!(
257 "A float initialiser was expected, found integers of shape {:?}.", dims;
258 Invalid, Input, Mismatch)),
259 }
260 }
261
262 /// Reads the initialiser as integers.
263 pub fn ints(&self) -> Outcome<&[i64]> {
264 match self {
265 Self::I64 { data, .. } => Ok(data),
266 Self::F32 { dims, .. } => Err(err!(
267 "An integer initialiser was expected, found floats of shape {:?}.", dims;
268 Invalid, Input, Mismatch)),
269 }
270 }
271}
272
273/// One node of the graph, as the model spelled it.
274#[derive(Clone, Debug, Default)]
275pub struct Node {
276 /// The operator name, such as `Conv`.
277 pub op: String,
278 /// The node's own name, which may be empty.
279 pub name: String,
280 /// Names of the tensors it consumes.
281 pub inputs: Vec<String>,
282 /// Names of the tensors it produces.
283 pub outputs: Vec<String>,
284 /// Attributes, in the order the model listed them.
285 pub attrs: Vec<(String, Attr)>,
286}
287
288impl Node {
289 /// Finds an attribute by name.
290 pub fn attr(&self, name: &str) -> Option<&Attr> {
291 self.attrs.iter().find(|(n, _)| n == name).map(|(_, a)| a)
292 }
293
294 /// Finds an attribute by name, failing if it is absent.
295 pub fn need(&self, name: &str) -> Outcome<&Attr> {
296 match self.attr(name) {
297 Some(a) => Ok(a),
298 None => Err(err!(
299 "The {} node has no {} attribute.", self.op, name; Invalid, Input, Missing)),
300 }
301 }
302}
303
304/// A model, read but not yet interpreted.
305#[derive(Clone, Debug, Default)]
306pub struct Model {
307 /// Nodes, in the order the model listed them, which ONNX requires to be
308 /// topological.
309 pub nodes: Vec<Node>,
310 /// Initialisers, by name.
311 pub inits: Vec<(String, Init)>,
312 /// Declared graph inputs. A model exported from some frameworks lists its
313 /// weights here as well, each shadowed by an initialiser of the same name.
314 pub inputs: Vec<String>,
315 /// Declared graph outputs, in order.
316 pub outputs: Vec<String>,
317}
318
319impl Model {
320 /// Finds an initialiser by name.
321 pub fn init(&self, name: &str) -> Option<&Init> {
322 self.inits.iter().find(|(n, _)| n == name).map(|(_, i)| i)
323 }
324
325 /// Reads a model from the bytes of an `.onnx` file.
326 pub fn read(bytes: &[u8]) -> Outcome<Self> {
327 let mut r = Reader::new(bytes);
328 let mut graph = None;
329 while !r.done() {
330 let (num, f) = res!(r.next());
331 if num == 7 {
332 if let Field::Bytes(b) = f {
333 graph = Some(b);
334 }
335 }
336 }
337 let g = match graph {
338 Some(g) => g,
339 None => return Err(err!("The model carries no graph."; Invalid, Input, Missing)),
340 };
341 Self::read_graph(g)
342 }
343
344 /// Reads a `GraphProto`.
345 fn read_graph(buf: &[u8]) -> Outcome<Self> {
346 let mut m = Self::default();
347 let mut r = Reader::new(buf);
348 while !r.done() {
349 let (num, f) = res!(r.next());
350 match (num, f) {
351 (1, Field::Bytes(b)) => m.nodes.push(res!(read_node(b))),
352 (5, Field::Bytes(b)) => {
353 let (name, init) = res!(read_tensor(b));
354 m.inits.push((name, init));
355 },
356 (11, Field::Bytes(b)) => m.inputs.push(res!(read_value_info(b))),
357 (12, Field::Bytes(b)) => m.outputs.push(res!(read_value_info(b))),
358 _ => {},
359 }
360 }
361 Ok(m)
362 }
363}
364
365/// Reads a `NodeProto`.
366fn read_node(buf: &[u8]) -> Outcome<Node> {
367 let mut n = Node::default();
368 let mut r = Reader::new(buf);
369 while !r.done() {
370 let (num, f) = res!(r.next());
371 match (num, f) {
372 (1, Field::Bytes(b)) => n.inputs.push(res!(as_str(b))),
373 (2, Field::Bytes(b)) => n.outputs.push(res!(as_str(b))),
374 (3, Field::Bytes(b)) => n.name = res!(as_str(b)),
375 (4, Field::Bytes(b)) => n.op = res!(as_str(b)),
376 (5, Field::Bytes(b)) => {
377 if let Some(a) = res!(read_attr(b)) {
378 n.attrs.push(a);
379 }
380 },
381 _ => {},
382 }
383 }
384 Ok(n)
385}
386
387/// Reads an `AttributeProto`, answering `None` for a kind this reader does not
388/// carry -- a subgraph, for instance, which no model here uses.
389fn read_attr(buf: &[u8]) -> Outcome<Option<(String, Attr)>> {
390 let mut name = String::new();
391 let mut typ = 0i64;
392 let mut i = 0i64;
393 let mut fl = 0f32;
394 let mut s = String::new();
395 let mut ints: Vec<i64> = Vec::new();
396 let mut floats: Vec<f32> = Vec::new();
397 let mut r = Reader::new(buf);
398 while !r.done() {
399 let (num, f) = res!(r.next());
400 match (num, f) {
401 (1, Field::Bytes(b)) => name = res!(as_str(b)),
402 (2, Field::Fixed32(a)) => fl = f32::from_le_bytes(a),
403 (3, Field::Varint(v)) => i = v as i64,
404 (4, Field::Bytes(b)) => s = res!(as_str(b)),
405 (7, Field::Bytes(b)) => floats = res!(le_f32(b)),
406 (7, Field::Fixed32(a)) => floats.push(f32::from_le_bytes(a)),
407 (8, Field::Bytes(b)) => ints = res!(packed_varints(b)),
408 (8, Field::Varint(v)) => ints.push(v as i64),
409 (20, Field::Varint(v)) => typ = v as i64,
410 _ => {},
411 }
412 }
413 // AttributeType: 1 FLOAT, 2 INT, 3 STRING, 6 FLOATS, 7 INTS.
414 let a = match typ {
415 1 => Attr::Float(fl),
416 2 => Attr::Int(i),
417 3 => Attr::Str(s),
418 6 => Attr::Floats(floats),
419 7 => Attr::Ints(ints),
420 _ => return Ok(None),
421 };
422 Ok(Some((name, a)))
423}
424
425/// Reads a `TensorProto`, answering its name and values.
426fn read_tensor(buf: &[u8]) -> Outcome<(String, Init)> {
427 let mut dims: Vec<usize> = Vec::new();
428 let mut dtype = 0i64;
429 let mut name = String::new();
430 let mut raw: Option<&[u8]> = None;
431 let mut floats: Vec<f32> = Vec::new();
432 let mut ints: Vec<i64> = Vec::new();
433 let mut r = Reader::new(buf);
434 while !r.done() {
435 let (num, f) = res!(r.next());
436 match (num, f) {
437 (1, Field::Varint(v)) => dims.push(v as usize),
438 (1, Field::Bytes(b)) => {
439 for v in res!(packed_varints(b)) {
440 dims.push(v as usize);
441 }
442 },
443 (2, Field::Varint(v)) => dtype = v as i64,
444 (4, Field::Bytes(b)) => floats = res!(le_f32(b)),
445 (7, Field::Bytes(b)) => ints = res!(packed_varints(b)),
446 (8, Field::Bytes(b)) => name = res!(as_str(b)),
447 (9, Field::Bytes(b)) => raw = Some(b),
448 _ => {},
449 }
450 }
451 let want = dims.iter().product::<usize>();
452 let init = match dtype {
453 DT_FLOAT => {
454 let data = match raw {
455 Some(b) => res!(le_f32(b)),
456 None => floats,
457 };
458 if data.len() != want {
459 return Err(err!(
460 "The initialiser {} declares shape {:?}, which wants {} floats, but \
461 carries {}.", name, dims, want, data.len();
462 Invalid, Input, Mismatch));
463 }
464 Init::F32 { dims, data }
465 },
466 DT_INT64 => {
467 let data = match raw {
468 Some(b) => res!(le_i64(b)),
469 None => ints,
470 };
471 if data.len() != want {
472 return Err(err!(
473 "The initialiser {} declares shape {:?}, which wants {} integers, but \
474 carries {}.", name, dims, want, data.len();
475 Invalid, Input, Mismatch));
476 }
477 Init::I64 { dims, data }
478 },
479 other => return Err(err!(
480 "The initialiser {} has element type {}, which this reader does not carry.",
481 name, other;
482 Invalid, Input, Unimplemented)),
483 };
484 Ok((name, init))
485}
486
487/// Reads the name out of a `ValueInfoProto`.
488fn read_value_info(buf: &[u8]) -> Outcome<String> {
489 let mut r = Reader::new(buf);
490 while !r.done() {
491 let (num, f) = res!(r.next());
492 if let (1, Field::Bytes(b)) = (num, f) {
493 return as_str(b);
494 }
495 }
496 Ok(String::new())
497}
498
499#[cfg(test)]
500mod tests {
501 use super::*;
502
503 #[test]
504 fn a_truncated_message_is_an_error_not_a_panic() -> Outcome<()> {
505 // A length-delimited field claiming more bytes than remain.
506 let bytes = [0x0a, 0x40, 0x01, 0x02];
507 let mut r = Reader::new(&bytes);
508 req!(r.next().is_err(), true);
509 Ok(())
510 }
511
512 #[test]
513 fn an_empty_model_has_no_graph() -> Outcome<()> {
514 req!(Model::read(&[]).is_err(), true);
515 Ok(())
516 }
517
518 #[test]
519 fn packed_dimensions_read_back() -> Outcome<()> {
520 let v = res!(packed_varints(&[0x01, 0x02, 0x80, 0x02]));
521 req!(v, vec![1i64, 2, 256]);
522 Ok(())
523 }
524}