Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/graph.rs

36.3 KiB, 52 runs

created by r1870400018:19740, 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//! The operator set, the weight preparation that happens once at load, and the
2//! runner that walks a prepared graph.
3//!
4//! # What preparation does
5//!
6//! A model as exported is not the shape a kernel wants. Loading therefore does
7//! four things, all once:
8//!
9//! - permutes every convolution weight from `[O, I, kh, kw]` into the
10//! `[kh·kw·I, O]` matrix a channels-last product consumes;
11//! - turns each `BatchNormalization` into a per-channel affine map, and folds
12//! that map into the convolution or matrix product in front of it wherever
13//! the intermediate value has no other reader;
14//! - drops the operators that are identities at inference -- dropout, and the
15//! transpose-and-reshape pair a detection head uses to relabel its output,
16//! which costs nothing once the activation is already channels-last;
17//! - resolves shapes that are written as an initialiser, such as a reshape
18//! target or a resize scale.
19//!
20//! What survives is a flat list of layers over a flat list of values.
21
22use crate::kern::{
23 self,
24 Coord,
25 Cpu,
26 Sample,
27 Scratch,
28 Task,
29};
30use crate::onnx;
31use crate::tensor::Tensor;
32
33use oxedyne_fe2o3_core::prelude::*;
34
35/// Everything a convolution needs, with its weights already permuted.
36#[derive(Clone, Debug)]
37pub struct Conv {
38 /// Output channels.
39 pub oc: usize,
40 /// Input channels. For a depthwise convolution this equals `oc`.
41 pub ic: usize,
42 /// Kernel height.
43 pub kh: usize,
44 /// Kernel width.
45 pub kw: usize,
46 /// Vertical stride.
47 pub sy: usize,
48 /// Horizontal stride.
49 pub sx: usize,
50 /// Padding above.
51 pub pt: usize,
52 /// Padding to the left.
53 pub pl: usize,
54 /// Padding below.
55 pub pb: usize,
56 /// Padding to the right.
57 pub pr: usize,
58 /// Weights: `[kh·kw·ic, oc]` when dense, `[kh·kw, oc]` when depthwise.
59 pub weight: Vec<f32>,
60 /// Per-output-channel bias, zero when the model gave none.
61 pub bias: Vec<f32>,
62}
63
64/// Everything a maximum pool needs.
65#[derive(Clone, Copy, Debug)]
66pub struct Pool {
67 /// Kernel height.
68 pub kh: usize,
69 /// Kernel width.
70 pub kw: usize,
71 /// Vertical stride.
72 pub sy: usize,
73 /// Horizontal stride.
74 pub sx: usize,
75 /// Padding above.
76 pub pt: usize,
77 /// Padding to the left.
78 pub pl: usize,
79 /// Padding below.
80 pub pb: usize,
81 /// Padding to the right.
82 pub pr: usize,
83}
84
85impl Pool {
86 /// The extent one axis pools down to.
87 pub fn out(&self, n: usize, k: usize, s: usize, before: usize, after: usize)
88 -> Outcome<usize>
89 {
90 let padded = n + before + after;
91 if padded < k {
92 return Err(err!(
93 "A pool of kernel {} wants at least that many samples, found {}.", k, padded;
94 Invalid, Input, Mismatch));
95 }
96 Ok((padded - k) / s + 1)
97 }
98}
99
100/// How large a resize makes its result.
101///
102/// Both forms occur in the models this crate carries: a face detector writes a
103/// scale, and a category detector writes the size it wants outright.
104#[derive(Clone, Copy, Debug)]
105pub enum Extent {
106 /// Multiply the height and the width.
107 Scale(f32, f32),
108 /// A fixed height and width.
109 Size(usize, usize),
110}
111
112impl Extent {
113 /// The output extents, given the input's.
114 pub fn applied(&self, h: usize, w: usize) -> Outcome<(usize, usize)> {
115 let (oh, ow) = match self {
116 Self::Scale(sy, sx) => (
117 (h as f32 * sy).round() as usize,
118 (w as f32 * sx).round() as usize,
119 ),
120 Self::Size(oh, ow) => (*oh, *ow),
121 };
122 if oh == 0 || ow == 0 {
123 return Err(err!(
124 "A resize of {} by {} to {} by {} has an empty axis.", h, w, oh, ow;
125 Invalid, Input, Range));
126 }
127 Ok((oh, ow))
128 }
129}
130
131/// One prepared operator.
132#[derive(Clone, Debug)]
133pub enum Op {
134 /// Dense convolution over every input channel.
135 Conv(Conv),
136 /// Depthwise convolution, one weight plane per channel.
137 Depthwise(Conv),
138 /// Per-channel affine map, which is what a batch normalisation becomes.
139 Scale {
140 /// Per-channel multiplier.
141 scale: Vec<f32>,
142 /// Per-channel offset.
143 bias: Vec<f32>,
144 },
145 /// Parametric rectified linear unit with a per-channel slope.
146 PRelu {
147 /// Per-channel negative slope.
148 slope: Vec<f32>,
149 },
150 /// Rectified linear unit.
151 Relu,
152 /// Leaky rectified linear unit, one slope for every channel.
153 Leaky(f32),
154 /// Cuts the channels into runs of the given widths, one output each.
155 Split(Vec<usize>),
156 /// Joins its operands along their channels, in the order given.
157 Concat,
158 /// Interleaves the channels of a grouped activation. See
159 /// [`kern::shuffle_channels`].
160 Shuffle {
161 /// How many groups the channels were split into.
162 groups: usize,
163 },
164 /// Logistic sigmoid.
165 Sigmoid,
166 /// Maximum pool over any kernel, stride and padding.
167 MaxPool(Pool),
168 /// Resampling of the two spatial axes, up or down.
169 Resize {
170 /// How large the result is.
171 extent: Extent,
172 /// How a value is taken.
173 sample: Sample,
174 /// How an output position maps back into the source.
175 coord: Coord,
176 },
177 /// Element-wise sum of two activations of the same shape.
178 Add,
179 /// Subtracts a scalar from every value.
180 SubScalar(f32),
181 /// Multiplies every value by a scalar.
182 MulScalar(f32),
183 /// Flattens a channels-last activation into the channels-first row a
184 /// matrix product downstream was trained against.
185 Flatten,
186 /// Matrix--vector product against weights stored `[n, k]`.
187 MatVecT {
188 /// Number of outputs.
189 n: usize,
190 /// Length of each dot product.
191 k: usize,
192 /// Weights, `[n, k]`.
193 weight: Vec<f32>,
194 /// Per-output bias.
195 bias: Vec<f32>,
196 },
197 /// Passes the value through unchanged.
198 Identity,
199 /// Rewrites the shape without moving a value.
200 Reshape(Vec<i64>),
201}
202
203/// One layer: an operator and the values it reads and writes.
204#[derive(Clone, Debug)]
205pub struct Layer {
206 /// What the layer computes.
207 pub op: Op,
208 /// Indices of the values it reads.
209 pub inputs: Vec<usize>,
210 /// Indices of the values it writes.
211 pub outputs: Vec<usize>,
212}
213
214/// A prepared graph, ready to run.
215#[derive(Clone, Debug)]
216pub struct Graph {
217 /// Layers, in execution order.
218 pub layers: Vec<Layer>,
219 /// Name of each value, for error messages.
220 pub names: Vec<String>,
221 /// Index of the value the caller supplies.
222 pub input: usize,
223 /// Indices of the values the caller wants back, in declared order.
224 pub outputs: Vec<usize>,
225}
226
227/// Interns tensor names as indices, so the runner needs no map.
228#[derive(Default)]
229struct Names {
230 /// Name of each value.
231 names: Vec<String>,
232}
233
234impl Names {
235 /// Answers the index of a name, adding it if it is new.
236 fn intern(&mut self, name: &str) -> usize {
237 match self.names.iter().position(|n| n == name) {
238 Some(i) => i,
239 None => {
240 self.names.push(name.to_string());
241 self.names.len() - 1
242 },
243 }
244 }
245}
246
247/// Reads a padding attribute, which ONNX writes as `[top, left, bottom, right]`.
248fn pads(node: &onnx::Node) -> Outcome<(usize, usize, usize, usize)> {
249 match node.attr("pads") {
250 None => Ok((0, 0, 0, 0)),
251 Some(a) => {
252 let v = res!(a.ints());
253 if v.len() != 4 {
254 return Err(err!(
255 "A two-dimensional convolution wants four padding values, found {}.", v.len();
256 Invalid, Input, Mismatch));
257 }
258 for p in v {
259 if *p < 0 {
260 return Err(err!("A negative padding of {} is not supported.", p;
261 Invalid, Input, Unimplemented));
262 }
263 }
264 Ok((v[0] as usize, v[1] as usize, v[2] as usize, v[3] as usize))
265 },
266 }
267}
268
269/// Reads a stride or dilation attribute, defaulting to one on each axis.
270fn pair(node: &onnx::Node, name: &str, default: usize) -> Outcome<(usize, usize)> {
271 match node.attr(name) {
272 None => Ok((default, default)),
273 Some(a) => {
274 let v = res!(a.ints());
275 if v.len() != 2 {
276 return Err(err!(
277 "The {} attribute wants two values, found {}.", name, v.len();
278 Invalid, Input, Mismatch));
279 }
280 if v[0] < 1 || v[1] < 1 {
281 return Err(err!("The {} attribute holds {:?}, which is not positive.", name, v;
282 Invalid, Input, Range));
283 }
284 Ok((v[0] as usize, v[1] as usize))
285 },
286 }
287}
288
289impl Graph {
290 /// Reads and prepares a model from the bytes of an `.onnx` file.
291 pub fn load(bytes: &[u8]) -> Outcome<Self> {
292 let m = res!(onnx::Model::read(bytes));
293 Self::prepare(&m)
294 }
295
296 /// Turns a read model into a prepared graph.
297 pub fn prepare(m: &onnx::Model) -> Outcome<Self> {
298 let mut names = Names::default();
299 let mut layers: Vec<Layer> = Vec::with_capacity(m.nodes.len());
300 let fixed = res!(relabellings(m));
301
302 for (ni, node) in m.nodes.iter().enumerate() {
303 if let Some(op) = fixed.get(&ni) {
304 let inputs = node.inputs.iter()
305 .filter(|n| m.init(n).is_none() && !n.is_empty())
306 .map(|n| names.intern(n))
307 .collect::<Vec<_>>();
308 let outputs = node.outputs.iter().map(|n| names.intern(n)).collect::<Vec<_>>();
309 layers.push(Layer { op: op.clone(), inputs, outputs });
310 continue;
311 }
312 let op = match node.op.as_str() {
313 "Conv" => res!(load_conv(m, node)),
314 "BatchNormalization" => res!(load_batch_norm(m, node)),
315 "PRelu" => res!(load_prelu(m, node)),
316 "Gemm" => res!(load_gemm(m, node)),
317 "Relu" => Op::Relu,
318 "LeakyRelu" => Op::Leaky(res!(res!(node.need("alpha")).float())),
319 "Split" => res!(load_split(node)),
320 "Concat" => res!(load_concat(node)),
321 "Sigmoid" => Op::Sigmoid,
322 "MaxPool" => res!(load_max_pool(node)),
323 "Resize" | "Upsample" => res!(load_resize(m, node)),
324 "Add" => Op::Add,
325 "Sub" => Op::SubScalar(res!(scalar_operand(m, node))),
326 "Mul" => Op::MulScalar(res!(scalar_operand(m, node))),
327 "Flatten" => Op::Flatten,
328 "Dropout" => Op::Identity,
329 "Transpose" => res!(load_transpose(node)),
330 "Reshape" => res!(load_reshape(m, node)),
331 other => return Err(err!(
332 "The operator {} is outside the subset this crate carries.", other;
333 Invalid, Input, Unimplemented)),
334 };
335 // Only the operands that are values, not the ones that were
336 // initialisers folded into the operator above.
337 let inputs = node.inputs.iter()
338 .filter(|n| m.init(n).is_none() && !n.is_empty())
339 .map(|n| names.intern(n))
340 .collect::<Vec<_>>();
341 let outputs = node.outputs.iter()
342 .map(|n| names.intern(n))
343 .collect::<Vec<_>>();
344 layers.push(Layer { op, inputs, outputs });
345 }
346
347 // The value the caller supplies is the first declared input with no
348 // initialiser shadowing it. A model exported from some frameworks lists
349 // every weight as an input as well.
350 let mut input = None;
351 for n in &m.inputs {
352 if m.init(n).is_none() {
353 input = Some(names.intern(n));
354 break;
355 }
356 }
357 let input = match input {
358 Some(i) => i,
359 None => return Err(err!(
360 "The graph declares no input that is not also an initialiser.";
361 Invalid, Input, Missing)),
362 };
363 let outputs = m.outputs.iter().map(|n| names.intern(n)).collect::<Vec<_>>();
364
365 let mut g = Self { layers, names: names.names, input, outputs };
366 res!(g.fold_scales());
367 Ok(g)
368 }
369
370 /// Folds each per-channel affine map into the operator that produced its
371 /// input, wherever that value has no other reader.
372 fn fold_scales(&mut self) -> Outcome<()> {
373 let mut i = 1;
374 while i < self.layers.len() {
375 let foldable = match (&self.layers[i - 1].op, &self.layers[i].op) {
376 (Op::Conv(_), Op::Scale { .. })
377 | (Op::Depthwise(_), Op::Scale { .. })
378 | (Op::MatVecT { .. }, Op::Scale { .. }) => {
379 let produced = self.layers[i - 1].outputs.first().copied();
380 let consumed = self.layers[i].inputs.first().copied();
381 produced.is_some()
382 && produced == consumed
383 && self.readers(produced.unwrap_or(usize::MAX)) == 1
384 && !self.outputs.contains(&produced.unwrap_or(usize::MAX))
385 },
386 _ => false,
387 };
388 if !foldable {
389 i += 1;
390 continue;
391 }
392 let (scale, bias) = match &self.layers[i].op {
393 Op::Scale { scale, bias } => (scale.clone(), bias.clone()),
394 _ => return Err(err!("A scale layer changed shape while folding.";
395 Bug, Invalid)),
396 };
397 let out = self.layers[i].outputs.clone();
398 match &mut self.layers[i - 1].op {
399 Op::Conv(c) | Op::Depthwise(c) => {
400 if scale.len() != c.oc {
401 return Err(err!(
402 "A scale of {} channels cannot fold into a convolution of {}.",
403 scale.len(), c.oc;
404 Invalid, Input, Mismatch));
405 }
406 let rows = c.weight.len() / c.oc;
407 for r in 0..rows {
408 for o in 0..c.oc {
409 c.weight[r * c.oc + o] *= scale[o];
410 }
411 }
412 for o in 0..c.oc {
413 c.bias[o] = c.bias[o] * scale[o] + bias[o];
414 }
415 },
416 Op::MatVecT { n, k, weight, bias: b } => {
417 if scale.len() != *n {
418 return Err(err!(
419 "A scale of {} channels cannot fold into a product of {} outputs.",
420 scale.len(), n;
421 Invalid, Input, Mismatch));
422 }
423 for o in 0..*n {
424 for p in 0..*k {
425 weight[o * *k + p] *= scale[o];
426 }
427 b[o] = b[o] * scale[o] + bias[o];
428 }
429 },
430 _ => return Err(err!("A layer changed shape while folding."; Bug, Invalid)),
431 }
432 self.layers[i - 1].outputs = out;
433 self.layers.remove(i);
434 }
435 Ok(())
436 }
437
438 /// Counts the layers that read a value.
439 fn readers(&self, value: usize) -> usize {
440 self.layers.iter().filter(|l| l.inputs.contains(&value)).count()
441 }
442
443 /// Runs the graph over one input, answering the declared outputs in order.
444 ///
445 /// The input is channels-last, `[N, H, W, C]`, and so is every activation
446 /// the runner passes between layers.
447 pub fn run(&self, cpu: Cpu, input: Tensor) -> Outcome<Vec<Tensor>> {
448 let mut left: Vec<usize> = vec![0; self.names.len()];
449 for l in &self.layers {
450 for v in &l.inputs {
451 left[*v] += 1;
452 }
453 }
454 // A declared output must survive to the end.
455 for v in &self.outputs {
456 left[*v] += 1;
457 }
458 let mut vals: Vec<Option<Tensor>> = vec![None; self.names.len()];
459 vals[self.input] = Some(input);
460 let mut scratch = Scratch::new();
461
462 for (li, l) in self.layers.iter().enumerate() {
463 let out = res!(self.step(cpu, l, &mut vals, &mut left, &mut scratch)
464 .map_err(|e| err!(e,
465 "Layer {} ({}) failed.", li, self.names.get(
466 l.outputs.first().copied().unwrap_or(0)).map(|s| s.as_str())
467 .unwrap_or("?");
468 Invalid)));
469 for (slot, t) in l.outputs.iter().zip(out.into_iter()) {
470 vals[*slot] = Some(t);
471 }
472 }
473
474 let mut answer = Vec::with_capacity(self.outputs.len());
475 for v in &self.outputs {
476 match vals[*v].take() {
477 Some(t) => answer.push(t),
478 None => return Err(err!(
479 "The graph output {} was never produced.", self.names[*v];
480 Invalid, Missing)),
481 }
482 }
483 Ok(answer)
484 }
485
486 /// Runs one layer.
487 fn step(
488 &self,
489 cpu: Cpu,
490 l: &Layer,
491 vals: &mut Vec<Option<Tensor>>,
492 left: &mut Vec<usize>,
493 scratch: &mut Scratch,
494 )
495 -> Outcome<Vec<Tensor>>
496 {
497 let x = res!(take(vals, left, l.inputs.first().copied(), &self.names));
498 let out = match &l.op {
499 Op::Conv(c) => res!(run_conv(cpu, c, &x, scratch, false)),
500 Op::Depthwise(c) => res!(run_conv(cpu, c, &x, scratch, true)),
501 Op::Scale { scale, bias } => {
502 let mut x = x;
503 let ch = *some!(x.dims.last(), "An activation with no axes cannot be scaled.");
504 kern::run(cpu, Task::Scale { ch, x: &mut x.data, scale, bias });
505 x
506 },
507 Op::PRelu { slope } => {
508 let mut x = x;
509 let ch = *some!(x.dims.last(), "An activation with no axes has no channels.");
510 kern::run(cpu, Task::PRelu { ch, x: &mut x.data, slope });
511 x
512 },
513 Op::Relu => {
514 let mut x = x;
515 kern::run(cpu, Task::Relu { x: &mut x.data });
516 x
517 },
518 Op::Sigmoid => {
519 let mut x = x;
520 kern::run(cpu, Task::Sigmoid { x: &mut x.data });
521 x
522 },
523 Op::MaxPool(p) => {
524 let (n, h, w, ch) = res!(x.nhwc());
525 let oh = res!(p.out(h, p.kh, p.sy, p.pt, p.pb));
526 let ow = res!(p.out(w, p.kw, p.sx, p.pl, p.pr));
527 let mut y = Tensor::zeros(vec![n, oh, ow, ch]);
528 kern::run(cpu, Task::MaxPool {
529 ch, h, w,
530 kh: p.kh, kw: p.kw, sy: p.sy, sx: p.sx, pt: p.pt, pl: p.pl,
531 oh, ow,
532 x: &x.data, y: &mut y.data,
533 });
534 y
535 },
536 Op::Resize { extent, sample, coord } => {
537 let (n, h, w, ch) = res!(x.nhwc());
538 let (oh, ow) = res!(extent.applied(h, w));
539 let mut y = Tensor::zeros(vec![n, oh, ow, ch]);
540 kern::run(cpu, Task::Resize {
541 ch, h, w, oh, ow,
542 sample: *sample, coord: *coord,
543 x: &x.data, y: &mut y.data,
544 });
545 y
546 },
547 Op::Add => {
548 let mut x = x;
549 let second = some!(l.inputs.get(1).copied(), "An addition wants two operands.");
550 let y = res!(take(vals, left, Some(second), &self.names));
551 if y.len() != x.len() {
552 return Err(err!(
553 "An addition of {:?} and {:?} does not line up.", x.dims, y.dims;
554 Invalid, Input, Mismatch));
555 }
556 kern::run(cpu, Task::Add { x: &mut x.data, y: &y.data });
557 x
558 },
559 Op::SubScalar(v) => {
560 let mut x = x;
561 for e in x.data.iter_mut() {
562 *e -= *v;
563 }
564 x
565 },
566 Op::MulScalar(v) => {
567 let mut x = x;
568 for e in x.data.iter_mut() {
569 *e *= *v;
570 }
571 x
572 },
573 Op::Flatten => res!(kern::flatten_nchw(&x)),
574 Op::MatVecT { n, k, weight, bias } => {
575 if x.len() != *k {
576 return Err(err!(
577 "A product of inner extent {} was given {} values.", k, x.len();
578 Invalid, Input, Mismatch));
579 }
580 let mut c = vec![0.0f32; *n];
581 kern::run(cpu, Task::MatVecT { n: *n, k: *k, a: &x.data, bt: weight, c: &mut c });
582 for (o, b) in c.iter_mut().zip(bias.iter()) {
583 *o += *b;
584 }
585 res!(Tensor::new(vec![1, *n], c))
586 },
587 Op::Leaky(slope) => {
588 let mut x = x;
589 kern::run(cpu, Task::Leaky { x: &mut x.data, slope: *slope });
590 x
591 },
592 Op::Split(widths) => return kern::split_channels(&x, widths),
593 Op::Concat => {
594 // The first operand came through `x`; the rest are still in the
595 // pool, and each is wanted only here.
596 let mut rest = Vec::with_capacity(l.inputs.len());
597 for idx in l.inputs.iter().skip(1) {
598 rest.push(res!(take(vals, left, Some(*idx), &self.names)));
599 }
600 let mut parts = Vec::with_capacity(rest.len() + 1);
601 parts.push(&x);
602 for t in &rest {
603 parts.push(t);
604 }
605 res!(kern::concat_channels(&parts))
606 },
607 Op::Shuffle { groups } => res!(kern::shuffle_channels(&x, *groups)),
608 Op::Identity => x,
609 Op::Reshape(spec) => {
610 let mut x = x;
611 let dims = res!(resolve_shape(spec, x.len(), &x.dims));
612 res!(x.reshape(dims));
613 x
614 },
615 };
616 Ok(vec![out])
617 }
618}
619
620/// Takes a value out of the pool, cloning only when another layer still wants it.
621fn take(
622 vals: &mut Vec<Option<Tensor>>,
623 left: &mut Vec<usize>,
624 idx: Option<usize>,
625 names: &[String],
626)
627 -> Outcome<Tensor>
628{
629 let i = some!(idx, "A layer names no input.");
630 if left[i] > 0 {
631 left[i] -= 1;
632 }
633 if left[i] == 0 {
634 match vals[i].take() {
635 Some(t) => Ok(t),
636 None => Err(err!("The value {} was read before it was written.", names[i];
637 Invalid, Missing)),
638 }
639 } else {
640 match &vals[i] {
641 Some(t) => Ok(t.clone()),
642 None => Err(err!("The value {} was read before it was written.", names[i];
643 Invalid, Missing)),
644 }
645 }
646}
647
648/// Resolves a reshape target, which may name a free axis with minus one and an
649/// axis to copy with zero.
650fn resolve_shape(spec: &[i64], total: usize, from: &[usize]) -> Outcome<Vec<usize>> {
651 let mut dims = Vec::with_capacity(spec.len());
652 let mut free = None;
653 let mut known = 1usize;
654 for (i, v) in spec.iter().enumerate() {
655 match *v {
656 -1 => {
657 if free.is_some() {
658 return Err(err!("A reshape names more than one free axis.";
659 Invalid, Input));
660 }
661 free = Some(i);
662 dims.push(1);
663 },
664 0 => {
665 let d = *some!(from.get(i), "A reshape copies an axis the input does not have.");
666 known *= d;
667 dims.push(d);
668 },
669 n if n > 0 => {
670 known *= n as usize;
671 dims.push(n as usize);
672 },
673 n => return Err(err!("A reshape names a negative extent of {}.", n;
674 Invalid, Input, Range)),
675 }
676 }
677 match free {
678 Some(i) => {
679 if known == 0 || total % known != 0 {
680 return Err(err!(
681 "A reshape of {} values into {:?} does not divide.", total, spec;
682 Invalid, Input, Mismatch));
683 }
684 dims[i] = total / known;
685 },
686 None => {
687 if known != total {
688 return Err(err!(
689 "A reshape of {} values into {:?} wants {}.", total, spec, known;
690 Invalid, Input, Mismatch));
691 }
692 },
693 }
694 Ok(dims)
695}
696
697/// Reads a convolution node, permuting its weights into the layout the kernels
698/// consume.
699fn load_conv(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
700 let wname = some!(node.inputs.get(1), "A convolution names no weight.");
701 let w = some!(m.init(wname), "The convolution weight is not an initialiser.");
702 let wd = w.dims();
703 if wd.len() != 4 {
704 return Err(err!(
705 "A two-dimensional convolution wants a weight of rank 4, found {:?}.", wd;
706 Invalid, Input, Mismatch));
707 }
708 let wv = res!(w.floats());
709 let (oc, icg, kh, kw) = (wd[0], wd[1], wd[2], wd[3]);
710 let group = match node.attr("group") {
711 Some(a) => res!(a.int()) as usize,
712 None => 1,
713 };
714 let (dy, dx) = res!(pair(node, "dilations", 1));
715 if dy != 1 || dx != 1 {
716 return Err(err!("A dilated convolution is outside the subset this crate carries.";
717 Invalid, Input, Unimplemented));
718 }
719 let (sy, sx) = res!(pair(node, "strides", 1));
720 let (pt, pl, pb, pr) = res!(pads(node));
721 let bias = match node.inputs.get(2) {
722 Some(bn) if !bn.is_empty() => {
723 let b = some!(m.init(bn), "The convolution bias is not an initialiser.");
724 res!(b.floats()).to_vec()
725 },
726 _ => vec![0.0; oc],
727 };
728 if bias.len() != oc {
729 return Err(err!(
730 "A convolution of {} outputs has a bias of {}.", oc, bias.len();
731 Invalid, Input, Mismatch));
732 }
733
734 if group == 1 {
735 // Dense: `[O, I, kh, kw]` becomes `[kh·kw·I, O]`.
736 let ic = icg;
737 let mut weight = vec![0.0f32; kh * kw * ic * oc];
738 for o in 0..oc {
739 for ci in 0..ic {
740 for ky in 0..kh {
741 for kx in 0..kw {
742 let src = ((o * ic + ci) * kh + ky) * kw + kx;
743 let dst = ((ky * kw + kx) * ic + ci) * oc + o;
744 weight[dst] = wv[src];
745 }
746 }
747 }
748 }
749 Ok(Op::Conv(Conv { oc, ic, kh, kw, sy, sx, pt, pl, pb, pr, weight, bias }))
750 } else if group == oc && icg == 1 {
751 // Depthwise: `[C, 1, kh, kw]` becomes `[kh·kw, C]`.
752 let mut weight = vec![0.0f32; kh * kw * oc];
753 for c in 0..oc {
754 for ky in 0..kh {
755 for kx in 0..kw {
756 weight[(ky * kw + kx) * oc + c] = wv[(c * kh + ky) * kw + kx];
757 }
758 }
759 }
760 Ok(Op::Depthwise(Conv {
761 oc,
762 ic: oc,
763 kh, kw, sy, sx, pt, pl, pb, pr, weight, bias,
764 }))
765 } else {
766 Err(err!(
767 "A grouped convolution with {} groups over {} input channels is outside the \
768 subset this crate carries.", group, icg * group;
769 Invalid, Input, Unimplemented))
770 }
771}
772
773/// Turns a batch normalisation into the per-channel affine map it is at
774/// inference.
775fn load_batch_norm(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
776 let eps = match node.attr("epsilon") {
777 Some(a) => res!(a.float()),
778 None => 1e-5,
779 };
780 let mut got = Vec::with_capacity(4);
781 for i in 1..5 {
782 let name = some!(node.inputs.get(i), "A batch normalisation wants four parameters.");
783 let init = some!(m.init(name), "A batch normalisation parameter is not an initialiser.");
784 got.push(res!(init.floats()));
785 }
786 let (gamma, beta, mean, var) = (got[0], got[1], got[2], got[3]);
787 let c = gamma.len();
788 if beta.len() != c || mean.len() != c || var.len() != c {
789 return Err(err!(
790 "A batch normalisation has parameters of {}, {}, {} and {} channels.",
791 c, beta.len(), mean.len(), var.len();
792 Invalid, Input, Mismatch));
793 }
794 let mut scale = vec![0.0f32; c];
795 let mut bias = vec![0.0f32; c];
796 for i in 0..c {
797 let denom = (var[i] + eps).sqrt();
798 if denom == 0.0 {
799 return Err(err!("A batch normalisation has a zero variance in channel {}.", i;
800 Invalid, Input, Range));
801 }
802 scale[i] = gamma[i] / denom;
803 bias[i] = beta[i] - mean[i] * scale[i];
804 }
805 Ok(Op::Scale { scale, bias })
806}
807
808/// Reads a parametric rectifier, whose slope the model stores as `[C, 1, 1]`.
809fn load_prelu(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
810 let name = some!(node.inputs.get(1), "A parametric rectifier names no slope.");
811 let init = some!(m.init(name), "The rectifier slope is not an initialiser.");
812 Ok(Op::PRelu { slope: res!(init.floats()).to_vec() })
813}
814
815/// Reads a matrix product. Only the transposed-weight form appears in the
816/// models this crate targets, and it is the form a contiguous dot product wants.
817fn load_gemm(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
818 let trans_a = match node.attr("transA") {
819 Some(a) => res!(a.int()),
820 None => 0,
821 };
822 let trans_b = match node.attr("transB") {
823 Some(a) => res!(a.int()),
824 None => 0,
825 };
826 let alpha = match node.attr("alpha") {
827 Some(a) => res!(a.float()),
828 None => 1.0,
829 };
830 let beta = match node.attr("beta") {
831 Some(a) => res!(a.float()),
832 None => 1.0,
833 };
834 if trans_a != 0 || trans_b != 1 {
835 return Err(err!(
836 "A matrix product with transA={} and transB={} is outside the subset this \
837 crate carries.", trans_a, trans_b;
838 Invalid, Input, Unimplemented));
839 }
840 let wname = some!(node.inputs.get(1), "A matrix product names no weight.");
841 let w = some!(m.init(wname), "The matrix product weight is not an initialiser.");
842 let wd = w.dims();
843 if wd.len() != 2 {
844 return Err(err!(
845 "A matrix product wants a weight of rank 2, found {:?}.", wd;
846 Invalid, Input, Mismatch));
847 }
848 let (n, k) = (wd[0], wd[1]);
849 let mut weight = res!(w.floats()).to_vec();
850 if alpha != 1.0 {
851 for v in weight.iter_mut() {
852 *v *= alpha;
853 }
854 }
855 let mut bias = match node.inputs.get(2) {
856 Some(bn) if !bn.is_empty() => {
857 let b = some!(m.init(bn), "The matrix product bias is not an initialiser.");
858 res!(b.floats()).to_vec()
859 },
860 _ => vec![0.0; n],
861 };
862 if beta != 1.0 {
863 for v in bias.iter_mut() {
864 *v *= beta;
865 }
866 }
867 if bias.len() != n {
868 return Err(err!(
869 "A matrix product of {} outputs has a bias of {}.", n, bias.len();
870 Invalid, Input, Mismatch));
871 }
872 Ok(Op::MatVecT { n, k, weight, bias })
873}
874
875/// Reads a maximum pool, which this crate carries only in its two by two,
876/// stride two, unpadded form.
877fn load_max_pool(node: &onnx::Node) -> Outcome<Op> {
878 let k = res!(res!(node.need("kernel_shape")).ints()).to_vec();
879 let (sy, sx) = res!(pair(node, "strides", 1));
880 let (pt, pl, pb, pr) = res!(pads(node));
881 let ceil = match node.attr("ceil_mode") {
882 Some(a) => res!(a.int()),
883 None => 0,
884 };
885 if k.len() != 2 {
886 return Err(err!(
887 "A two-dimensional maximum pool wants a kernel of two extents, found {:?}.", k;
888 Invalid, Input, Mismatch));
889 }
890 // Rounding the output up rather than down would need the kernel to run off
891 // the end of the source, which nothing here does.
892 if ceil != 0 {
893 return Err(err!(
894 "A maximum pool rounding its output up is outside the subset this crate carries.";
895 Invalid, Input, Unimplemented));
896 }
897 Ok(Op::MaxPool(Pool {
898 kh: k[0] as usize,
899 kw: k[1] as usize,
900 sy, sx, pt, pl, pb, pr,
901 }))
902}
903
904/// Reads a resize of the two spatial axes.
905///
906/// ONNX writes the target either as a scale or as an outright size, and the
907/// convention mapping an output position back into the source as an attribute.
908/// All three are read rather than assumed: the two models this crate carries
909/// disagree on every one of them, and a half-sample error in the mapping moves
910/// every box a detector predicts.
911fn load_resize(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
912 let sample = match node.attr("mode") {
913 None => Sample::Nearest,
914 Some(a) => match res!(a.text()) {
915 "nearest" => Sample::Nearest,
916 "linear" => Sample::Bilinear,
917 other => return Err(err!(
918 "A resize in {} mode is outside the subset this crate carries, which is \
919 nearest and linear.", other;
920 Invalid, Input, Unimplemented)),
921 },
922 };
923 let coord = match node.attr("coordinate_transformation_mode") {
924 None => Coord::Asymmetric,
925 Some(a) => match res!(a.text()) {
926 "asymmetric" => Coord::Asymmetric,
927 "pytorch_half_pixel" => Coord::HalfPixel,
928 // `half_pixel` differs from PyTorch's only for an output of one,
929 // which no model here asks for, but it is not silently accepted.
930 other => return Err(err!(
931 "A resize transforming coordinates by {} is outside the subset this crate \
932 carries.", other;
933 Invalid, Input, Unimplemented)),
934 },
935 };
936
937 // The target is the last constant operand: floats are scales, integers are
938 // the size outright. An empty initialiser is the `roi` operand, skipped.
939 let mut extent: Option<Extent> = None;
940 for name in node.inputs.iter().skip(1) {
941 let init = match m.init(name) {
942 Some(i) => i,
943 None => continue,
944 };
945 if let Ok(v) = init.floats() {
946 if v.len() == 4 {
947 extent = Some(Extent::Scale(v[2], v[3]));
948 continue;
949 }
950 }
951 if let Ok(v) = init.ints() {
952 if v.len() == 4 {
953 extent = Some(Extent::Size(v[2] as usize, v[3] as usize));
954 }
955 }
956 }
957 let extent = some!(extent, "A resize names neither a scale nor a size operand.");
958 Ok(Op::Resize { extent, sample, coord })
959}
960
961/// Reads the scalar second operand of an element-wise node.
962fn scalar_operand(m: &onnx::Model, node: &onnx::Node) -> Outcome<f32> {
963 let name = some!(node.inputs.get(1), "An element-wise node names no second operand.");
964 let init = some!(m.init(name), "The second operand is not an initialiser.");
965 let v = res!(init.floats());
966 if v.len() != 1 {
967 return Err(err!(
968 "An element-wise node with a second operand of {} values is outside the subset \
969 this crate carries, which is a scalar.", v.len();
970 Invalid, Input, Unimplemented));
971 }
972 Ok(v[0])
973}
974
975/// Reads a transpose. Channels-first to channels-last is what the activations
976/// already are, so it costs nothing; anything else would need a real permutation.
977///
978/// The three-axis form is the same relabelling after the spatial axes have been
979/// folded into one, which is how a detection head presents its predictions.
980fn load_transpose(node: &onnx::Node) -> Outcome<Op> {
981 let perm = res!(res!(node.need("perm")).ints()).to_vec();
982 if perm == vec![0, 2, 3, 1] || perm == vec![0, 2, 1] {
983 Ok(Op::Identity)
984 } else {
985 Err(err!(
986 "A transpose by {:?} is outside the subset this crate carries, which is the \
987 channels-first to channels-last relabelling.", perm;
988 Invalid, Input, Unimplemented))
989 }
990}
991
992/// Reads a split, which this crate carries along the channels only.
993fn load_split(node: &onnx::Node) -> Outcome<Op> {
994 let axis = match node.attr("axis") {
995 Some(a) => res!(a.int()),
996 None => 0,
997 };
998 if axis != 1 {
999 return Err(err!(
1000 "A split along axis {} is outside the subset this crate carries, which is the \
1001 channels.", axis;
1002 Invalid, Input, Unimplemented));
1003 }
1004 let widths = res!(res!(node.need("split")).ints());
1005 let mut out = Vec::with_capacity(widths.len());
1006 for w in widths {
1007 if *w <= 0 {
1008 return Err(err!("A split of width {} is not a run of channels.", w;
1009 Invalid, Input, Range));
1010 }
1011 out.push(*w as usize);
1012 }
1013 Ok(Op::Split(out))
1014}
1015
1016/// Reads a concatenation, which this crate carries along the channels only.
1017fn load_concat(node: &onnx::Node) -> Outcome<Op> {
1018 let axis = match node.attr("axis") {
1019 Some(a) => res!(a.int()),
1020 None => 1,
1021 };
1022 if axis != 1 {
1023 return Err(err!(
1024 "A concatenation along axis {} is outside the subset this crate carries, which \
1025 is the channels.", axis;
1026 Invalid, Input, Unimplemented));
1027 }
1028 Ok(Op::Concat)
1029}
1030
1031/// Finds the runs of nodes that are shape bookkeeping in a channels-first model
1032/// and nothing at all in a channels-last one, and says what each becomes.
1033///
1034/// Two shapes occur, and both are written as a reshape with a transpose in the
1035/// middle because that is the only way a channels-first graph can say them.
1036///
1037/// - **A channel shuffle**, `[n, c, h, w]` to `[n, g, c/g, h, w]`, transposed by
1038/// `[0, 2, 1, 3, 4]` and reshaped back. Here the spatial axes never move and
1039/// the whole of it is a permutation of the innermost axis, so the two reshapes
1040/// are nothing and the transpose is the permutation.
1041/// - **A detection head's relabelling**, `[n, c, h, w]` to `[n, c, h·w]`
1042/// transposed by `[0, 2, 1]`. Channels-last data is already in the order the
1043/// result wants, so the pair together is one reshape that moves no value.
1044///
1045/// Matching the run rather than the node is deliberate. Each of these operators
1046/// alone would need a real permutation of a channels-last activation; it is only
1047/// as a run that they come to nothing, and a graph that used one on its own
1048/// should be refused rather than quietly mishandled.
1049fn relabellings(m: &onnx::Model) -> Outcome<std::collections::BTreeMap<usize, Op>> {
1050 let mut out = std::collections::BTreeMap::new();
1051 for (i, node) in m.nodes.iter().enumerate() {
1052 if node.op != "Transpose" {
1053 continue;
1054 }
1055 let perm = match node.attr("perm") {
1056 Some(a) => res!(a.ints()).to_vec(),
1057 None => continue,
1058 };
1059
1060 if perm == vec![0, 2, 1, 3, 4] {
1061 // The reshape before it names the groups in its second position.
1062 if i == 0 || m.nodes[i - 1].op != "Reshape" || i + 1 >= m.nodes.len()
1063 || m.nodes[i + 1].op != "Reshape"
1064 {
1065 return Err(err!(
1066 "A five-axis transpose by {:?} is only carried as the middle of a \
1067 channel shuffle, and this one is not.", perm;
1068 Invalid, Input, Unimplemented));
1069 }
1070 let target = res!(reshape_target(m, &m.nodes[i - 1]));
1071 if target.len() != 5 || target[1] <= 0 {
1072 return Err(err!(
1073 "A channel shuffle reshapes to {:?}, which names no group count.", target;
1074 Invalid, Input, Mismatch));
1075 }
1076 out.insert(i - 1, Op::Identity);
1077 out.insert(i, Op::Shuffle { groups: target[1] as usize });
1078 out.insert(i + 1, Op::Identity);
1079 } else if perm == vec![0, 2, 1] {
1080 // The reshape before it says how many channels the head predicts.
1081 if i == 0 || m.nodes[i - 1].op != "Reshape" {
1082 continue;
1083 }
1084 let target = res!(reshape_target(m, &m.nodes[i - 1]));
1085 if target.len() != 3 {
1086 continue;
1087 }
1088 out.insert(i - 1, Op::Reshape(vec![target[0], -1, target[1]]));
1089 out.insert(i, Op::Identity);
1090 }
1091 }
1092 Ok(out)
1093}
1094
1095/// The target shape a reshape names, which the model stores as an initialiser.
1096fn reshape_target(m: &onnx::Model, node: &onnx::Node) -> Outcome<Vec<i64>> {
1097 let name = some!(node.inputs.get(1), "A reshape names no target.");
1098 let init = some!(m.init(name), "The reshape target is not an initialiser.");
1099 Ok(res!(init.ints()).to_vec())
1100}
1101
1102/// Reads a reshape, whose target the model stores as an initialiser.
1103fn load_reshape(m: &onnx::Model, node: &onnx::Node) -> Outcome<Op> {
1104 let name = some!(node.inputs.get(1), "A reshape names no target.");
1105 let init = some!(m.init(name), "The reshape target is not an initialiser.");
1106 Ok(Op::Reshape(res!(init.ints()).to_vec()))
1107}
1108
1109/// Runs one convolution, gathering patches first when the kernel is larger than
1110/// one by one.
1111fn run_conv(
1112 cpu: Cpu,
1113 c: &Conv,
1114 x: &Tensor,
1115 scratch: &mut Scratch,
1116 depthwise: bool,
1117)
1118 -> Outcome<Tensor>
1119{
1120 let (n, h, w, ch) = res!(x.nhwc());
1121 if n != 1 {
1122 return Err(err!("A batch of {} is outside the subset this crate carries.", n;
1123 Invalid, Input, Unimplemented));
1124 }
1125 if ch != c.ic {
1126 return Err(err!(
1127 "A convolution over {} input channels was given an activation of {}.", c.ic, ch;
1128 Invalid, Input, Mismatch));
1129 }
1130 let oh = res!(out_extent(h, c.kh, c.sy, c.pt, c.pb));
1131 let ow = res!(out_extent(w, c.kw, c.sx, c.pl, c.pr));
1132 let mut y = Tensor::zeros(vec![1, oh, ow, c.oc]);
1133
1134 if depthwise {
1135 kern::run(cpu, Task::Depthwise {
1136 ch: c.oc,
1137 h, w,
1138 kh: c.kh,
1139 kw: c.kw,
1140 sy: c.sy,
1141 sx: c.sx,
1142 pt: c.pt,
1143 pl: c.pl,
1144 oh, ow,
1145 x: &x.data,
1146 wt: &c.weight,
1147 bias: Some(&c.bias),
1148 y: &mut y.data,
1149 });
1150 return Ok(y);
1151 }
1152
1153 let unit = c.kh == 1 && c.kw == 1 && c.sy == 1 && c.sx == 1
1154 && c.pt == 0 && c.pl == 0 && c.pb == 0 && c.pr == 0;
1155 if unit {
1156 // In this layout the activation already *is* the patch matrix.
1157 kern::run(cpu, Task::Gemm {
1158 m: oh * ow,
1159 n: c.oc,
1160 k: c.ic,
1161 a: &x.data,
1162 b: &c.weight,
1163 c: &mut y.data,
1164 bias: Some(&c.bias),
1165 scratch,
1166 });
1167 return Ok(y);
1168 }
1169
1170 let k = c.kh * c.kw * c.ic;
1171 let mut patches = vec![0.0f32; oh * ow * k];
1172 kern::run(cpu, Task::Im2Col {
1173 ch,
1174 h, w,
1175 kh: c.kh,
1176 kw: c.kw,
1177 sy: c.sy,
1178 sx: c.sx,
1179 pt: c.pt,
1180 pl: c.pl,
1181 oh, ow,
1182 x: &x.data,
1183 out: &mut patches,
1184 });
1185 kern::run(cpu, Task::Gemm {
1186 m: oh * ow,
1187 n: c.oc,
1188 k,
1189 a: &patches,
1190 b: &c.weight,
1191 c: &mut y.data,
1192 bias: Some(&c.bias),
1193 scratch,
1194 });
1195 Ok(y)
1196}
1197
1198/// The extent of one output axis of a convolution.
1199fn out_extent(n: usize, k: usize, stride: usize, before: usize, after: usize)
1200 -> Outcome<usize>
1201{
1202 let padded = n + before + after;
1203 if padded < k {
1204 return Err(err!(
1205 "A kernel of {} does not fit an axis of {} padded to {}.", k, n, padded;
1206 Invalid, Input, Range));
1207 }
1208 Ok((padded - k) / stride + 1)
1209}
1210
1211#[cfg(test)]
1212mod tests {
1213 use super::*;
1214
1215 #[test]
1216 fn a_free_axis_is_resolved() -> Outcome<()> {
1217 req!(res!(resolve_shape(&[1, -1, 4], 48, &[1, 12, 4])), vec![1, 12, 4]);
1218 req!(res!(resolve_shape(&[0, 6], 48, &[8, 6])), vec![8, 6]);
1219 req!(resolve_shape(&[1, -1, -1], 48, &[1, 12, 4]).is_err(), true);
1220 req!(resolve_shape(&[5, 5], 48, &[1]).is_err(), true);
1221 Ok(())
1222 }
1223
1224 #[test]
1225 fn an_output_extent_follows_the_padding() -> Outcome<()> {
1226 req!(res!(out_extent(640, 3, 2, 1, 1)), 320);
1227 req!(res!(out_extent(112, 3, 1, 1, 1)), 112);
1228 req!(res!(out_extent(112, 1, 1, 0, 0)), 112);
1229 req!(out_extent(1, 3, 1, 0, 0).is_err(), true);
1230 Ok(())
1231 }
1232}