Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/tensor.rs

3.9 KiB, 1 run

created by r1870400018:19737, 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 activation and weight container the graph runner passes between operators.
2
3use oxedyne_fe2o3_core::prelude::*;
4
5/// A dense `f32` tensor with row-major, contiguous data.
6///
7/// Four-dimensional activations are held as `[N, H, W, C]` -- channels last --
8/// which is the layout every kernel in this crate expects. An ONNX model
9/// declares its activations as `[N, C, H, W]`, so the loader permutes the
10/// weights once and the runner never transposes an activation again.
11#[derive(Clone, Debug, Default, PartialEq)]
12pub struct Tensor {
13 /// Extent along each axis, outermost first.
14 pub dims: Vec<usize>,
15 /// Values, row-major over `dims`.
16 pub data: Vec<f32>,
17}
18
19impl Tensor {
20 /// Creates a tensor from dimensions and values, checking that they agree.
21 pub fn new(dims: Vec<usize>, data: Vec<f32>) -> Outcome<Self> {
22 let want = dims.iter().product::<usize>();
23 if want != data.len() {
24 return Err(err!(
25 "A tensor of shape {:?} holds {} values, but {} were given.",
26 dims, want, data.len();
27 Invalid, Input, Mismatch));
28 }
29 Ok(Self { dims, data })
30 }
31
32 /// Creates a zeroed tensor of the given shape.
33 pub fn zeros(dims: Vec<usize>) -> Self {
34 let n = dims.iter().product::<usize>();
35 Self { dims, data: vec![0.0; n] }
36 }
37
38 /// Number of values in the tensor.
39 pub fn len(&self) -> usize {
40 self.data.len()
41 }
42
43 /// Whether the tensor holds no values.
44 pub fn is_empty(&self) -> bool {
45 self.data.is_empty()
46 }
47
48 /// Number of axes.
49 pub fn rank(&self) -> usize {
50 self.dims.len()
51 }
52
53 /// Reads the tensor as a four-dimensional `[N, H, W, C]` activation.
54 pub fn nhwc(&self) -> Outcome<(usize, usize, usize, usize)> {
55 if self.dims.len() != 4 {
56 return Err(err!(
57 "An activation of rank 4 was expected, found shape {:?}.", self.dims;
58 Invalid, Input, Mismatch));
59 }
60 Ok((self.dims[0], self.dims[1], self.dims[2], self.dims[3]))
61 }
62
63 /// Rewrites the shape, keeping the values, and checking the element count.
64 pub fn reshape(&mut self, dims: Vec<usize>) -> Outcome<()> {
65 let want = dims.iter().product::<usize>();
66 if want != self.data.len() {
67 return Err(err!(
68 "A reshape to {:?} wants {} values, but the tensor holds {}.",
69 dims, want, self.data.len();
70 Invalid, Input, Mismatch));
71 }
72 self.dims = dims;
73 Ok(())
74 }
75
76 /// Converts an `[N, C, H, W]` tensor to the `[N, H, W, C]` layout the
77 /// kernels use.
78 pub fn nchw_to_nhwc(&self) -> Outcome<Self> {
79 if self.dims.len() != 4 {
80 return Err(err!(
81 "A tensor of rank 4 was expected, found shape {:?}.", self.dims;
82 Invalid, Input, Mismatch));
83 }
84 let (n, c, h, w) = (self.dims[0], self.dims[1], self.dims[2], self.dims[3]);
85 let mut out = vec![0.0f32; self.data.len()];
86 for bi in 0..n {
87 for ci in 0..c {
88 let src = (bi * c + ci) * h * w;
89 for p in 0..h * w {
90 out[(bi * h * w + p) * c + ci] = self.data[src + p];
91 }
92 }
93 }
94 Ok(Self { dims: vec![n, h, w, c], data: out })
95 }
96
97 /// Converts an `[N, H, W, C]` tensor back to the `[N, C, H, W]` layout an
98 /// ONNX graph declares, which is what an external comparison wants.
99 pub fn nhwc_to_nchw(&self) -> Outcome<Self> {
100 if self.dims.len() != 4 {
101 return Err(err!(
102 "A tensor of rank 4 was expected, found shape {:?}.", self.dims;
103 Invalid, Input, Mismatch));
104 }
105 let (n, h, w, c) = (self.dims[0], self.dims[1], self.dims[2], self.dims[3]);
106 let mut out = vec![0.0f32; self.data.len()];
107 for bi in 0..n {
108 for ci in 0..c {
109 let dst = (bi * c + ci) * h * w;
110 for p in 0..h * w {
111 out[dst + p] = self.data[(bi * h * w + p) * c + ci];
112 }
113 }
114 }
115 Ok(Self { dims: vec![n, c, h, w], data: out })
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn layout_round_trip() -> Outcome<()> {
125 let t = res!(Tensor::new(
126 vec![1, 2, 2, 3],
127 (0..12).map(|v| v as f32).collect(),
128 ));
129 let nhwc = res!(t.nchw_to_nhwc());
130 req!(nhwc.dims, vec![1, 2, 3, 2]);
131 let back = res!(nhwc.nhwc_to_nchw());
132 req!(back, t);
133 Ok(())
134 }
135}