载入中...
搜索中...
未找到
OnnxExecute.cpp
浏览该文件的文档.
1#include "tensor/CpuKernels.h"
3#include "tensor/Tensor.h"
4
5#include <algorithm>
6#include <cmath>
7#include <numeric>
8#include <set>
9
11bool isSupported(const Node& n) {
12 if (n.domain == "com.microsoft" && n.op == "DynamicQuantizeLSTM") {
13 if (n.inputs.size() != 12 || n.outputs.empty() || n.outputs.size() > 3) return false;
14 for (const auto& [key, a] : n.attrs) {
15 if (key == "direction") {
16 if (a.type != 3 || (a.text != "forward" && a.text != "reverse" && a.text != "bidirectional"))
17 return false;
18 } else if (key == "hidden_size" || key == "input_forget") {
19 if (a.type != 2) return false;
20 } else if (key == "clip") {
21 if (a.type != 1) return false;
22 } else
23 return false;
24 }
25 return true;
26 }
27 if (!n.domain.empty()) return false;
28 if (n.op == "If" || n.op == "Loop") {
29 const std::set<std::string> keys =
30 n.op == "If" ? std::set<std::string>{"then_branch", "else_branch"} : std::set<std::string>{"body"};
31 if (n.attrs.size() != keys.size() || n.outputs.empty() ||
32 (n.op == "If" ? n.inputs.size() != 1 : n.inputs.size() < 2))
33 return false;
34 for (const auto& key : keys) {
35 auto it = n.attrs.find(key);
36 if (it == n.attrs.end() || it->second.type != 5 || !it->second.graph) return false;
37 }
38 return true;
39 }
40
41 static const std::map<std::string, std::set<std::string>> allowed{
42 {"DynamicQuantizeLinear", {}},
43 {"QuantizeLinear", {"axis"}},
44 {"DequantizeLinear", {"axis"}},
45 {"MatMulInteger", {}},
46 {"ConvInteger", {"auto_pad", "dilations", "group", "kernel_shape", "pads", "strides"}},
47 {"Constant", {"value"}},
48 {"Clip", {}},
49 {"Pad", {"mode"}},
50 {"CumSum", {"exclusive", "reverse"}},
51 {"SequenceEmpty", {"dtype"}},
52 {"SequenceAt", {}},
53 {"SequenceInsert", {}},
54 {"SplitToSequence", {"axis", "keepdims"}},
55 {"ConcatFromSequence", {"axis", "new_axis"}},
56 {"RandomUniformLike", {"dtype", "seed", "low", "high"}},
57 {"RandomNormalLike", {"dtype", "seed", "mean", "scale"}},
58 {"InstanceNormalization", {"epsilon"}},
59 {"ConvTranspose", {"auto_pad", "dilations", "group", "kernel_shape", "output_padding", "pads", "strides"}},
60 {"Resize", {"coordinate_transformation_mode", "mode", "nearest_mode", "cubic_coeff_a"}},
61 {"Range", {}},
62 {"TopK", {"axis", "largest", "sorted"}},
63 {"ScatterND", {"reduction"}},
64 {"ScatterElements", {"axis", "reduction"}},
65 {"ReduceProd", {"axes", "keepdims"}},
66 {"ReduceMax", {"axes", "keepdims"}},
67 {"Slice", {}},
68 {"Concat", {"axis"}},
69 {"Expand", {}},
70 {"ConstantOfShape", {"value"}},
71 {"Where", {}},
72 {"Equal", {}},
73 {"Less", {}},
74 {"Greater", {}},
75 {"And", {}},
76 {"Not", {}},
77 {"Reciprocal", {}},
78 {"LeakyRelu", {"alpha"}},
79 {"Floor", {}},
80 {"Round", {}},
81 {"Atan", {}},
82 {"Identity", {}},
83 {"Cast", {"to"}},
84 {"Gather", {"axis"}},
85 {"Shape", {"start", "end"}},
86 {"Reshape", {"allowzero"}},
87 {"Transpose", {"perm"}},
88 {"Unsqueeze", {}},
89 {"Squeeze", {}},
90 {"MatMul", {}},
91 {"Add", {}},
92 {"Sub", {}},
93 {"Mul", {}},
94 {"Div", {}},
95 {"Sqrt", {}},
96 {"Exp", {}},
97 {"Log", {}},
98 {"Sin", {}},
99 {"Cos", {}},
100 {"Relu", {}},
101 {"Sigmoid", {}},
102 {"Tanh", {}},
103 {"Neg", {}},
104 {"Abs", {}},
105 {"Pow", {}},
106 {"Softmax", {"axis"}},
107 {"ReduceMean", {"axes", "keepdims"}},
108 {"ReduceSum", {"keepdims", "noop_with_empty_axes"}}};
109 auto op = allowed.find(n.op);
110 if (op == allowed.end()) return false;
111 size_t minimum = 1, maximum = 1, outputs = 1;
112 if (n.op == "Clip") maximum = 3;
113 if (n.op == "Pad") {
114 minimum = 2;
115 maximum = 3;
116 }
117 if (n.op == "CumSum") minimum = maximum = 2;
118 if (n.op == "SequenceEmpty") minimum = maximum = 0;
119 if (n.op == "SequenceAt") minimum = maximum = 2;
120 if (n.op == "SequenceInsert") {
121 minimum = 2;
122 maximum = 3;
123 }
124 if (n.op == "SplitToSequence") maximum = 2;
125 if (n.op == "InstanceNormalization") minimum = maximum = 3;
126 if (n.op == "ConvTranspose") {
127 minimum = 2;
128 maximum = 3;
129 }
130 if (n.op == "Resize") {
131 minimum = 3;
132 maximum = 4;
133 }
134 if (n.op == "Range" || n.op == "ScatterElements" || n.op == "ScatterND") minimum = maximum = 3;
135 if (n.op == "TopK") {
136 minimum = maximum = 2;
137 outputs = 2;
138 }
139 if (n.op == "Slice") {
140 minimum = 3;
141 maximum = 5;
142 }
143 if (n.op == "Concat") maximum = 100000;
144 if (n.op == "Where") minimum = maximum = 3;
145 if (n.op == "Expand" || n.op == "Equal" || n.op == "Less" || n.op == "Greater" || n.op == "And")
146 minimum = maximum = 2;
147 if (n.op == "Constant") minimum = maximum = 0;
148 if (n.op == "DynamicQuantizeLinear") outputs = 3;
149 if (n.op == "QuantizeLinear" || n.op == "DequantizeLinear") {
150 minimum = 2;
151 maximum = 3;
152 }
153 if (n.op == "MatMulInteger" || n.op == "ConvInteger") {
154 minimum = 2;
155 maximum = 4;
156 }
157 if (n.op == "Gather" || n.op == "Reshape" || n.op == "Unsqueeze" || n.op == "MatMul" || n.op == "Add" ||
158 n.op == "Sub" || n.op == "Mul" || n.op == "Div" || n.op == "Pow")
159 minimum = maximum = 2;
160 if (n.op == "Squeeze" || n.op == "ReduceSum") maximum = 2;
161 if (n.inputs.size() < minimum || n.inputs.size() > maximum || n.outputs.size() != outputs) return false;
162 for (const auto& [key, a] : n.attrs) {
163 if (!op->second.contains(key)) return false;
164 const int expected = (key == "alpha" || key == "epsilon" || key == "cubic_coeff_a" || key == "seed" ||
165 key == "low" || key == "high" || key == "mean" || key == "scale")
166 ? 1
167 : key == "value" ? 4
168 : (key == "auto_pad" || key == "reduction" || key == "coordinate_transformation_mode" ||
169 key == "mode" || key == "nearest_mode")
170 ? 3
171 : (key == "perm" || key == "axes" || key == "dilations" || key == "kernel_shape" ||
172 key == "pads" || key == "strides" || key == "output_padding")
173 ? 7
174 : 2;
175 if (a.type != expected) return false;
176 if (key == "reduction" && a.text != "none") return false;
177 if (key == "auto_pad" && a.text != "NOTSET") return false;
178 if (key == "allowzero" && a.integer != 0) return false;
179 }
180 if (n.op == "Pad" && n.attrs.contains("mode") && n.attrs.at("mode").text != "constant" &&
181 n.attrs.at("mode").text != "edge" && n.attrs.at("mode").text != "reflect")
182 return false;
183 if (n.op == "Resize") {
184 const auto mode = n.attrs.contains("mode") ? n.attrs.at("mode").text : "nearest";
185 const auto coordinate = n.attrs.contains("coordinate_transformation_mode")
186 ? n.attrs.at("coordinate_transformation_mode").text
187 : "half_pixel";
188 const auto nearest = n.attrs.contains("nearest_mode") ? n.attrs.at("nearest_mode").text : "round_prefer_floor";
189 if (mode == "nearest") {
190 if (coordinate != "asymmetric" || nearest != "floor") return false;
191 } else if (mode != "linear" || (coordinate != "half_pixel" && coordinate != "asymmetric"))
192 return false;
193 }
194 return true;
195}
196namespace {
197std::vector<int> dims(const RuntimeTensor& x) {
198 std::vector<int> out;
199 for (auto d : x.shape) {
200 if (d <= 0) throw Failure("Empty float kernel input", DiagnosticCode::Unsupported);
201 out.push_back(static_cast<int>(d));
202 }
203 return out;
204}
205std::vector<RuntimeTensor> single(RuntimeTensor x) { return {std::move(x)}; }
206} // namespace
207std::vector<RuntimeTensor> execute(const Node& n, const std::vector<const RuntimeTensor*>& in, OnnxCompute* compute) {
208 if (n.domain == "com.microsoft" && n.op == "DynamicQuantizeLSTM") return executeQuantLstm(n, in, compute);
209 if (n.op == "Constant") {
210 auto a = n.attrs.find("value");
211 if (a == n.attrs.end() || !a->second.tensor) throw Failure("Constant requires tensor value");
212 return single(*a->second.tensor);
213 }
214 if (n.op.find("Quantize") != std::string::npos || n.op == "DequantizeLinear" || n.op == "MatMulInteger" ||
215 n.op == "ConvInteger")
216 return executeQuant(n, in, compute);
217 const auto& x = required(in, 0);
218 if (n.op == "Identity") return single(x);
219 if (n.op == "Shape") {
220 const int64_t rank = x.shape.size();
221 auto norm = [rank](int64_t i) { return std::clamp(i < 0 ? i + rank : i, int64_t(0), rank); };
222 const auto begin = norm(attr(n, "start", 0)), end = norm(attr(n, "end", rank));
223 std::vector<int64_t> shape;
224 if (begin < end) shape.assign(x.shape.begin() + begin, x.shape.begin() + end);
225 return single(make(OnnxElement::Int64, {static_cast<int64_t>(shape.size())}, shape));
226 }
227 if (n.op == "Cast") {
228 const auto target = static_cast<OnnxElement>(attr(n, "to", 0));
230 return single(dispatchFloat(*compute, {&x}, x.shape, "y[i]=float(x0[i]);"));
231 const auto size = elementSize(target);
232 RuntimeTensor out{target, x.shape, std::vector<uint8_t>(count(x.shape) * size)};
233 for (size_t i = 0; i < count(x.shape); ++i) {
235 const float v =
236 x.element == OnnxElement::Float32 ? read<float>(x, i) : static_cast<float>(integer(x, i));
237 std::memcpy(out.bytes.data() + i * 4, &v, 4);
238 } else {
239 int64_t v;
240 if (x.element == OnnxElement::Float32) {
241 const double f = read<float>(x, i);
242 if (!std::isfinite(f) || f < -9223372036854775808.0 || f >= 9223372036854775808.0)
243 throw Failure("Cast overflow");
244 v = target == OnnxElement::Bool ? (f != 0) : static_cast<int64_t>(f);
245 } else
246 v = integer(x, i);
247 if (target == OnnxElement::Bool) v = v != 0;
248 std::memcpy(out.bytes.data() + i * size, &v, size);
249 }
250 }
251 return single(std::move(out));
252 }
253 if (n.op == "Gather") {
254 const auto& idx = required(in, 1);
255 const int a = axis(attr(n, "axis", 0), x.shape.size());
256 if (idx.element != OnnxElement::Int32 && idx.element != OnnxElement::Int64)
257 throw Failure("Gather indices must be int32/int64");
258 size_t inner = 1, outer = 1;
259 for (size_t i = a + 1; i < x.shape.size(); ++i) inner *= x.shape[i];
260 for (int i = 0; i < a; ++i) outer *= x.shape[i];
261 std::vector<int64_t> shape(x.shape.begin(), x.shape.begin() + a);
262 shape.insert(shape.end(), idx.shape.begin(), idx.shape.end());
263 shape.insert(shape.end(), x.shape.begin() + a + 1, x.shape.end());
264 const size_t stride = inner * elementSize(x.element), ni = count(idx.shape);
265 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
266 for (size_t o = 0; o < outer; ++o)
267 for (size_t i = 0; i < ni; ++i) {
268 auto j = integer(idx, i);
269 if (j < 0) j += x.shape[a];
270 if (j < 0 || j >= x.shape[a]) throw Failure("Gather index out of range");
271 if (stride)
272 std::memcpy(out.bytes.data() + (o * ni + i) * stride,
273 x.bytes.data() + (o * x.shape[a] + j) * stride, stride);
274 }
275 return single(std::move(out));
276 }
277 if (n.op == "Reshape" || n.op == "Unsqueeze" || n.op == "Squeeze") {
278 auto shape = x.shape;
279 if (n.op == "Reshape") {
280 shape = ints(required(in, 1));
281 int infer = -1;
282 size_t known = 1;
283 for (size_t i = 0; i < shape.size(); ++i) {
284 if (shape[i] == 0) {
285 if (i >= x.shape.size()) throw Failure("Reshape zero axis out of range");
286 shape[i] = x.shape[i];
287 }
288 if (shape[i] == -1) {
289 if (infer != -1) throw Failure("Multiple inferred dimensions");
290 infer = static_cast<int>(i);
291 } else {
292 if (shape[i] < 0 || shape[i] > INT32_MAX ||
293 (shape[i] && known > 128u * 1024u * 1024u / static_cast<size_t>(shape[i])))
294 throw Failure("Reshape overflow");
295 known *= static_cast<size_t>(shape[i]);
296 }
297 }
298 if (infer >= 0) {
299 if (!known || count(x.shape) % known) throw Failure("Invalid inferred shape");
300 shape[infer] = count(x.shape) / known;
301 }
302 } else {
303 std::vector<int64_t> axes;
304 if (in.size() > 1 && in[1])
305 axes = ints(*in[1]);
306 else if (n.op == "Unsqueeze")
307 throw Failure("Unsqueeze needs axes");
308 else
309 for (size_t i = 0; i < shape.size(); ++i)
310 if (shape[i] == 1) axes.push_back(i);
311 const size_t rank = shape.size() + (n.op == "Unsqueeze" ? axes.size() : 0);
312 std::set<int> positions;
313 for (auto a : axes)
314 if (!positions.insert(axis(a, rank)).second) throw Failure("Duplicate squeeze axis");
315 if (n.op == "Unsqueeze") {
316 for (int a : positions) shape.insert(shape.begin() + a, 1);
317 } else
318 for (auto it = positions.rbegin(); it != positions.rend(); ++it) {
319 if (shape[*it] != 1) throw Failure("Squeeze dimension not one");
320 shape.erase(shape.begin() + *it);
321 }
322 }
323 if (count(shape) != count(x.shape)) throw Failure("Reshape element count mismatch");
324 RuntimeTensor out = x;
325 out.shape = std::move(shape);
326 return single(std::move(out));
327 }
328 if (n.op == "Transpose") {
329 std::vector<int64_t> order = attrs(n, "perm", {});
330 if (order.empty())
331 for (size_t i = x.shape.size(); i > 0; --i) order.push_back(i - 1);
332 if (order.size() != x.shape.size()) throw Failure("Invalid transpose rank");
333 std::set<int64_t> used;
334 std::vector<int64_t> shape;
335 for (auto a : order) {
336 if (a < 0 || a >= static_cast<int64_t>(x.shape.size()) || !used.insert(a).second)
337 throw Failure("Invalid transpose permutation");
338 shape.push_back(x.shape[a]);
339 }
340 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(x.bytes.size())};
341 const auto size = elementSize(x.element);
342 std::vector<size_t> strides(x.shape.size(), 1);
343 for (size_t i = x.shape.size(); i > 1; --i) strides[i - 2] = strides[i - 1] * x.shape[i - 1];
344 for (size_t i = 0; i < count(shape); ++i) {
345 size_t rest = i, source = 0;
346 for (size_t j = shape.size(); j > 0; --j) {
347 source += (rest % shape[j - 1]) * strides[order[j - 1]];
348 rest /= shape[j - 1];
349 }
350 std::memcpy(out.bytes.data() + i * size, x.bytes.data() + source * size, size);
351 }
352 return single(std::move(out));
353 }
354 if (auto indexResult = executeIndex(n, in)) return std::move(*indexResult);
355 if (auto shapeResult = executeShape(n, in)) return single(std::move(*shapeResult));
356 if (auto miscResult = executeMisc(n, in, compute)) return single(std::move(*miscResult));
357 if (auto neuralResult = executeNeural(n, in, compute)) return single(std::move(*neuralResult));
358 if (auto numericResult = executeNumeric(n, in, compute)) return single(std::move(*numericResult));
359 const auto a = floats(x);
360 if (n.op == "Softmax") {
361 auto xd = dims(x);
362 std::vector<float> out(a.size());
363 kernels::softmax(a.data(), xd.data(), static_cast<int>(xd.size()), axis(attr(n, "axis", -1), xd.size()), false,
364 out.data());
365 return single(make(OnnxElement::Float32, x.shape, out));
366 }
367 static const std::map<std::string, OpType> unary{
368 {"Sqrt", OpType::Sqrt}, {"Exp", OpType::Exp}, {"Log", OpType::Log}, {"Sin", OpType::Sin},
369 {"Cos", OpType::Cos}, {"Relu", OpType::Relu}, {"Sigmoid", OpType::Sigmoid}, {"Tanh", OpType::Tanh},
370 {"Neg", OpType::Neg}, {"Abs", OpType::Abs}};
371 auto op = unary.find(n.op);
372 if (op == unary.end()) throw Failure("Unsupported operator", DiagnosticCode::Unsupported);
373 std::vector<float> out(a.size());
374 kernels::unaryOp(op->second, a.data(), static_cast<int>(a.size()), out.data(), 0, 0);
375 return single(make(OnnxElement::Float32, x.shape, out));
376}
377} // namespace eve::tensor::onnx_detail
LogicalId target
float x
Definition AnimClip.cpp:738
float maximum[3]
float minimum[3]
std::uint32_t key
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
std::vector< float > positions
float v
bool required
MeleePoint3 a
Definition MeleeHit.cpp:40
std::vector< std::int32_t > order
OnnxCompute * compute
int idx
float f
float begin
float d
std::uint32_t count
float size
Definition TreeMesh.cpp:156
const UnitySourceAsset & source
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
void softmax(const float *in, const int *dims, int rank, int axis, bool logMode, float *out)
Softmax.
void unaryOp(OpType type, const float *in, int count, float *out, float s0, float s1)
Unary op.
int64_t integer(const RuntimeTensor &v, size_t i=0)
Integer.
std::optional< RuntimeTensor > executeMisc(const Node &, const std::vector< const RuntimeTensor * > &, OnnxCompute *)
Execute misc.
Definition OnnxMisc.cpp:7
std::vector< int64_t > ints(const RuntimeTensor &v)
Ints.
int64_t attr(const Node &n, const char *key, int64_t fallback)
Attr.
std::optional< RuntimeTensor > executeNumeric(const Node &, const std::vector< const RuntimeTensor * > &, OnnxCompute *)
Execute numeric.
std::vector< int64_t > attrs(const Node &n, const char *key, std::vector< int64_t > fallback)
Attrs.
std::vector< RuntimeTensor > executeQuant(const Node &node, const std::vector< const RuntimeTensor * > &inputs, OnnxCompute *compute=nullptr)
Execute quant.
Definition OnnxQuant.cpp:28
std::optional< RuntimeTensor > executeNeural(const Node &, const std::vector< const RuntimeTensor * > &, OnnxCompute *)
Execute neural.
Definition OnnxNeural.cpp:7
RuntimeTensor make(OnnxElement e, std::vector< int64_t > shape, const std::vector< T > &values)
Make.
int axis(int64_t a, size_t rank)
Axis.
std::vector< float > floats(const RuntimeTensor &v)
Floats.
RuntimeTensor dispatchFloat(OnnxCompute &, const std::vector< const RuntimeTensor * > &, const std::vector< int64_t > &, const std::string &, size_t work=0)
Dispatches float.
bool isSupported(const Node &n)
True when supported.
size_t elementSize(OnnxElement e)
Element size.
std::vector< RuntimeTensor > executeQuantLstm(const Node &node, const std::vector< const RuntimeTensor * > &inputs, OnnxCompute *compute=nullptr)
Execute quant lstm.
Definition OnnxLstm.cpp:78
std::optional< RuntimeTensor > executeShape(const Node &, const std::vector< const RuntimeTensor * > &)
Execute shape.
std::vector< RuntimeTensor > execute(const Node &n, const std::vector< const RuntimeTensor * > &in, OnnxCompute *compute)
Execute.
std::optional< std::vector< RuntimeTensor > > executeIndex(const Node &n, const std::vector< const RuntimeTensor * > &in)
Execute index.
OnnxElement
ONNX wire element types; distinct from block-quantized Tensor storage.
Definition OnnxModel.h:18