载入中...
搜索中...
未找到
OnnxInternal.h
浏览该文件的文档.
1#pragma once
2#include <cstring>
3#include <map>
4#include <optional>
5#include <stdexcept>
6#include <unordered_map>
8#include "tensor/OnnxModel.h"
9
12struct Failure : std::runtime_error {
14 std::string path;
16 explicit Failure(std::string message, DiagnosticCode c = DiagnosticCode::InvalidArgument, std::string p = {})
18 : std::runtime_error(std::move(message)), code(c), path(std::move(p)) {}
19};
20struct ModelData;
22struct Attribute {
23 int type = 0;
24 int64_t integer = 0;
25 float real = 0;
26 std::string text;
27 std::vector<int64_t> integers;
28 std::vector<float> reals;
29 std::vector<std::string> strings;
30 std::optional<RuntimeTensor> tensor;
31 std::shared_ptr<ModelData> graph;
32};
34struct Node {
35 std::string name, op, domain;
36 std::vector<std::string> inputs, outputs;
37 std::map<std::string, Attribute> attrs;
38};
40struct Input {
41 int type = 0;
42 std::vector<int64_t> shape;
43};
45struct ModelData {
47 std::vector<Node> nodes;
48 std::unordered_map<std::string, RuntimeTensor> constants;
49 std::map<std::string, Input> inputs;
50 std::map<std::string, int64_t> opsets;
51};
53inline size_t elementSize(OnnxElement e) {
54 switch (e) {
56 case OnnxElement::Int32: return 4;
57 case OnnxElement::Int64: return 8;
60 case OnnxElement::Bool: return 1;
61 }
63 throw Failure("Unsupported ONNX element type", DiagnosticCode::Unsupported);
64}
66inline size_t count(const std::vector<int64_t>& shape) {
67 if (shape.size() > 6) throw Failure("ONNX rank exceeds six", DiagnosticCode::Unsupported);
68 size_t n = 1;
69 for (int64_t d : shape) {
70 if (d < 0 || d > INT32_MAX || (d && n > (128u * 1024u * 1024u) / static_cast<size_t>(d)))
72 throw Failure("Invalid or excessive tensor shape");
73 n *= static_cast<size_t>(d);
74 }
75 return n;
76}
78inline void validate(const RuntimeTensor& v) {
79 if (count(v.shape) * elementSize(v.element) != v.bytes.size()) throw Failure("Tensor byte count mismatch");
80}
81template <class T>
83T read(const RuntimeTensor& v, size_t i) {
84 T x;
86 std::memcpy(&x, v.bytes.data() + i * sizeof(T), sizeof(T));
87 return x;
88}
89template <class T>
91RuntimeTensor make(OnnxElement e, std::vector<int64_t> shape, const std::vector<T>& values) {
92 RuntimeTensor v{e, std::move(shape), std::vector<uint8_t>(values.size() * sizeof(T))};
93 if (!values.empty()) std::memcpy(v.bytes.data(), values.data(), v.bytes.size());
95 validate(v);
96 return v;
97}
99inline int64_t integer(const RuntimeTensor& v, size_t i = 0) {
100 if (i >= count(v.shape)) throw Failure("Integer index out of bounds");
101 switch (v.element) {
102 case OnnxElement::Int64: return read<int64_t>(v, i);
103 case OnnxElement::Int32: return read<int32_t>(v, i);
104 case OnnxElement::Int8: return v.bytes[i] >= 128 ? int(v.bytes[i]) - 256 : v.bytes[i];
106 case OnnxElement::Bool: return v.bytes[i];
108 default: throw Failure("Expected integer tensor");
109 }
110}
112inline std::vector<float> floats(const RuntimeTensor& v) {
113 if (v.element != OnnxElement::Float32) throw Failure("Expected FP32 tensor");
115 std::vector<float> out(count(v.shape));
116 if (!out.empty()) std::memcpy(out.data(), v.bytes.data(), v.bytes.size());
117 return out;
118}
120inline std::vector<int64_t> ints(const RuntimeTensor& v) {
122 std::vector<int64_t> out(count(v.shape));
123 for (size_t i = 0; i < out.size(); ++i) out[i] = integer(v, i);
124 return out;
125}
127inline int64_t attr(const Node& n, const char* key, int64_t fallback) {
128 auto it = n.attrs.find(key);
129 return it == n.attrs.end() ? fallback : it->second.integer;
130}
132inline std::vector<int64_t> attrs(const Node& n, const char* key, std::vector<int64_t> fallback) {
133 auto it = n.attrs.find(key);
134 return it == n.attrs.end() ? fallback : it->second.integers;
135}
137inline const RuntimeTensor& required(const std::vector<const RuntimeTensor*>& in, size_t i) {
138 if (i >= in.size() || !in[i]) throw Failure("Missing required input");
139 return *in[i];
140}
142inline int axis(int64_t a, size_t rank) {
143 if (a < 0) a += rank;
144 if (a < 0 || a >= static_cast<int64_t>(rank)) throw Failure("Invalid axis");
145 return static_cast<int>(a);
146}
148std::vector<int64_t> broadcast(const std::vector<int64_t>& a, const std::vector<int64_t>& b);
150size_t broadcastIndex(size_t i, const std::vector<int64_t>& shape, const std::vector<int64_t>& output);
152std::optional<std::vector<RuntimeTensor>> executeIndex(const Node&, const std::vector<const RuntimeTensor*>&);
154std::optional<RuntimeTensor> executeShape(const Node&, const std::vector<const RuntimeTensor*>&);
156RuntimeTensor dispatchFloat(OnnxCompute&, const std::vector<const RuntimeTensor*>&, const std::vector<int64_t>&,
157 const std::string&, size_t work = 0);
159std::optional<RuntimeTensor> executeMisc(const Node&, const std::vector<const RuntimeTensor*>&, OnnxCompute*);
161std::optional<RuntimeTensor> executeNeural(const Node&, const std::vector<const RuntimeTensor*>&, OnnxCompute*);
163std::optional<RuntimeTensor> executeNumeric(const Node&, const std::vector<const RuntimeTensor*>&, OnnxCompute*);
165std::vector<OnnxNamedTensor> evaluate(const ModelData&, std::span<const OnnxNamedTensor>,
166 const std::vector<std::string>&, OnnxCompute*, OnnxRunOptions options);
168ModelData parse(std::span<const uint8_t> bytes);
170bool isSupported(const Node& node);
172std::vector<RuntimeTensor> execute(const Node& node, const std::vector<const RuntimeTensor*>& inputs,
173 OnnxCompute* compute = nullptr);
175std::vector<RuntimeTensor> executeQuant(const Node& node, const std::vector<const RuntimeTensor*>& inputs,
176 OnnxCompute* compute = nullptr);
178std::vector<RuntimeTensor> executeQuantLstm(const Node& node, const std::vector<const RuntimeTensor*>& inputs,
179 OnnxCompute* compute = nullptr);
180} // namespace eve::tensor::onnx_detail
float x
Definition AnimClip.cpp:738
std::string output
glm::vec4 p[6]
std::map< std::string, Var > values
std::string message
std::uint32_t key
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
int inputs
Definition GridGraph.cpp:23
float v
std::int32_t c
bool required
std::uint64_t bytes
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
OnnxCompute * compute
float d
const RoadNode * node
std::uint32_t count
const SquirrelValueOptions & options
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
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.
size_t broadcastIndex(size_t i, const std::vector< int64_t > &shape, const std::vector< int64_t > &output)
Broadcast index.
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.
std::vector< OnnxNamedTensor > evaluate(const ModelData &, std::span< const OnnxNamedTensor >, const std::vector< std::string > &, OnnxCompute *, OnnxRunOptions options)
Evaluate.
bool isSupported(const Node &n)
True when supported.
void validate(const RuntimeTensor &v)
Validate.
size_t elementSize(OnnxElement e)
Element size.
ModelData parse(std::span< const uint8_t > bytes)
Parse.
std::vector< RuntimeTensor > executeQuantLstm(const Node &node, const std::vector< const RuntimeTensor * > &inputs, OnnxCompute *compute=nullptr)
Execute quant lstm.
Definition OnnxLstm.cpp:78
std::vector< int64_t > broadcast(const std::vector< int64_t > &a, const std::vector< int64_t > &b)
Broadcast.
std::optional< RuntimeTensor > executeShape(const Node &, const std::vector< const RuntimeTensor * > &)
Execute shape.
T read(const RuntimeTensor &v, size_t i)
Reads read.
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
DiagnosticCode
Stable machine-readable diagnostic codes.
Definition Diagnostic.h:47
Owning admission report; unsupported nodes remain inspectable but cannot execute.
Definition OnnxModel.h:48
Per-call deterministic RNG and optional strict finite-output diagnostic.
Definition OnnxModel.h:37
std::optional< RuntimeTensor > tensor
std::vector< std::string > strings
std::shared_ptr< ModelData > graph
std::vector< int64_t > integers
Failure(std::string message, DiagnosticCode c=DiagnosticCode::InvalidArgument, std::string p={})
Constructs a Failure.
std::vector< int64_t > shape
std::map< std::string, int64_t > opsets
std::map< std::string, Input > inputs
std::unordered_map< std::string, RuntimeTensor > constants
std::vector< std::string > inputs
std::map< std::string, Attribute > attrs
std::vector< std::string > outputs