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 | |
| 12 | use oxedyne_fe2o3_core::prelude::*; |
| 13 | |
| 14 | /// Protocol buffer wire types this reader understands. |
| 15 | const WIRE_VARINT: u8 = 0; |
| 16 | const WIRE_I64: u8 = 1; |
| 17 | const WIRE_LEN: u8 = 2; |
| 18 | const WIRE_I32: u8 = 5; |
| 19 | |
| 20 | /// ONNX tensor element types this reader understands. |
| 21 | const DT_FLOAT: i64 = 1; |
| 22 | const DT_INT64: i64 = 7; |
| 23 | |
| 24 | /// A cursor over a protocol buffer message. |
| 25 | struct 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. |
| 33 | enum 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 | |
| 44 | impl<'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. |
| 124 | fn 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. |
| 132 | fn 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. |
| 142 | fn 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. |
| 156 | fn 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)] |
| 173 | pub 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 | |
| 186 | impl 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)] |
| 226 | pub 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 | |
| 243 | impl 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)] |
| 275 | pub 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 | |
| 288 | impl 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)] |
| 306 | pub 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 | |
| 319 | impl 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`. |
| 366 | fn 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. |
| 389 | fn 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. |
| 426 | fn 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`. |
| 488 | fn 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)] |
| 500 | mod 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 | } |