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 | |
| 22 | use crate::kern::{ |
| 23 | self, |
| 24 | Coord, |
| 25 | Cpu, |
| 26 | Sample, |
| 27 | Scratch, |
| 28 | Task, |
| 29 | }; |
| 30 | use crate::onnx; |
| 31 | use crate::tensor::Tensor; |
| 32 | |
| 33 | use oxedyne_fe2o3_core::prelude::*; |
| 34 | |
| 35 | /// Everything a convolution needs, with its weights already permuted. |
| 36 | #[derive(Clone, Debug)] |
| 37 | pub 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)] |
| 66 | pub 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 | |
| 85 | impl 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)] |
| 105 | pub 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 | |
| 112 | impl 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)] |
| 133 | pub 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)] |
| 205 | pub 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)] |
| 216 | pub 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)] |
| 229 | struct Names { |
| 230 | /// Name of each value. |
| 231 | names: Vec<String>, |
| 232 | } |
| 233 | |
| 234 | impl 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]`. |
| 248 | fn 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. |
| 270 | fn 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 | |
| 289 | impl 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. |
| 621 | fn 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. |
| 650 | fn 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. |
| 699 | fn 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. |
| 775 | fn 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]`. |
| 809 | fn 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. |
| 817 | fn 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. |
| 877 | fn 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. |
| 911 | fn 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. |
| 962 | fn 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. |
| 980 | fn 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. |
| 993 | fn 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. |
| 1017 | fn 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. |
| 1049 | fn 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. |
| 1096 | fn 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. |
| 1103 | fn 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. |
| 1111 | fn 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. |
| 1199 | fn 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)] |
| 1212 | mod 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 | } |