载入中...
搜索中...
未找到
OnnxReader.cpp
浏览该文件的文档.
2
3#include <bit>
4#include <functional>
5#include <set>
6
8namespace {
9struct Field {
10 int id = 0, wire = 0;
11 uint64_t number = 0;
12 std::span<const uint8_t> bytes;
13 std::string text() const { return {reinterpret_cast<const char*>(bytes.data()), bytes.size()}; }
14};
15class Reader {
16public:
17 explicit Reader(std::span<const uint8_t> data) : data_(data) {}
18 bool empty() const { return offset_ == data_.size(); }
19 uint64_t varint() {
20 uint64_t n = 0;
21 for (int shift = 0; shift < 70; shift += 7) {
22 if (offset_ == data_.size()) throw Failure("Truncated protobuf varint", DiagnosticCode::ParseError);
23 const uint8_t b = data_[offset_++];
24 if (shift == 63 && b > 1) throw Failure("Overflowing protobuf varint", DiagnosticCode::ParseError);
25 n |= static_cast<uint64_t>(b & 127) << shift;
26 if (!(b & 128)) return n;
27 }
28 throw Failure("Invalid protobuf varint", DiagnosticCode::ParseError);
29 }
30 std::span<const uint8_t> take(uint64_t n) {
31 if (n > data_.size() - offset_) throw Failure("Truncated protobuf field", DiagnosticCode::ParseError);
32 auto result = data_.subspan(offset_, static_cast<size_t>(n));
33 offset_ += static_cast<size_t>(n);
34 return result;
35 }
36 Field next() {
37 const uint64_t tag = varint();
38 if ((tag >> 3) == 0 || (tag >> 3) > 0x1fffffff)
39 throw Failure("Invalid protobuf tag", DiagnosticCode::ParseError);
40 Field f;
41 f.id = static_cast<int>(tag >> 3);
42 f.wire = tag & 7;
43 switch (f.wire) {
44 case 0: f.number = varint(); break;
45 case 1: f.bytes = take(8); break;
46 case 2: f.bytes = take(varint()); break;
47 case 5: f.bytes = take(4); break;
48 default: throw Failure("Unsupported protobuf wire type", DiagnosticCode::ParseError);
49 }
50 return f;
51 }
52
53private:
54 std::span<const uint8_t> data_;
55 size_t offset_ = 0;
56};
57void wire(const Field& f, int expected) {
58 if (f.wire != expected) throw Failure("Unexpected protobuf field type", DiagnosticCode::ParseError);
59}
60void integers(const Field& f, std::vector<int64_t>& out) {
61 if (f.wire == 0)
62 out.push_back(std::bit_cast<int64_t>(f.number));
63 else {
64 wire(f, 2);
65 Reader r(f.bytes);
66 while (!r.empty()) out.push_back(std::bit_cast<int64_t>(r.varint()));
67 }
68}
69void reals(const Field& f, std::vector<float>& out) {
70 if (f.wire != 2 && f.wire != 5) throw Failure("Invalid float wire type", DiagnosticCode::ParseError);
71 if (f.bytes.size() % 4) throw Failure("Truncated float payload", DiagnosticCode::ParseError);
72 for (size_t i = 0; i < f.bytes.size(); i += 4) {
73 float x;
74 std::memcpy(&x, f.bytes.data() + i, 4);
75 out.push_back(x);
76 }
77}
78std::pair<std::string, RuntimeTensor> tensor(std::span<const uint8_t> data) {
79 Reader r(data);
80 RuntimeTensor v;
81 std::string name;
82 std::vector<int64_t> iv;
83 std::vector<float> fv;
84 bool raw = false, type = false;
85 while (!r.empty()) {
86 auto f = r.next();
87 switch (f.id) {
88 case 1: integers(f, v.shape); break;
89 case 2:
90 wire(f, 0);
91 if (f.number > 9) throw Failure("Unsupported tensor element type", DiagnosticCode::Unsupported);
92 v.element = static_cast<OnnxElement>(f.number);
93 type = true;
94 break;
95 case 4: reals(f, fv); break;
96 case 5:
97 case 7: integers(f, iv); break;
98 case 8:
99 wire(f, 2);
100 name = f.text();
101 break;
102 case 9:
103 wire(f, 2);
104 if (raw) throw Failure("Duplicate raw tensor data", DiagnosticCode::ParseError);
105 v.bytes.assign(f.bytes.begin(), f.bytes.end());
106 raw = true;
107 break;
108 case 14:
109 wire(f, 0);
110 if (f.number != 0) throw Failure("External tensor data unsupported", DiagnosticCode::Unsupported);
111 break;
112 case 3:
113 case 6:
114 case 10:
115 case 11:
116 case 13: throw Failure("External, segmented or unsupported tensor data", DiagnosticCode::Unsupported);
117 default: break;
118 }
119 }
120 if (!type) throw Failure("Tensor has no dtype", DiagnosticCode::ParseError);
121 const size_t n = count(v.shape), size = elementSize(v.element);
122 if (raw && (!iv.empty() || !fv.empty())) throw Failure("Ambiguous tensor payload", DiagnosticCode::ParseError);
123 if (!raw) {
124 if (v.element == OnnxElement::Float32) {
125 if (fv.size() != n || !iv.empty())
126 throw Failure("Invalid float tensor payload", DiagnosticCode::ParseError);
127 v.bytes.resize(n * size);
128 if (n) std::memcpy(v.bytes.data(), fv.data(), v.bytes.size());
129 } else {
130 if (iv.size() != n || !fv.empty())
131 throw Failure("Invalid integer tensor payload", DiagnosticCode::ParseError);
132 v.bytes.resize(n * size);
133 for (size_t i = 0; i < n; ++i) {
134 const auto x = iv[i];
135 if ((v.element == OnnxElement::UInt8 && (x < 0 || x > 255)) ||
136 (v.element == OnnxElement::Int8 && (x < -128 || x > 127)) ||
137 (v.element == OnnxElement::Bool && x != 0 && x != 1) ||
138 (v.element == OnnxElement::Int32 && (x < INT32_MIN || x > INT32_MAX)))
139 throw Failure("Integer initializer out of range", DiagnosticCode::ParseError);
140 std::memcpy(v.bytes.data() + i * size, &x, size);
141 }
142 }
143 }
144 validate(v);
145 return {std::move(name), std::move(v)};
146}
147ModelData parseGraph(std::span<const uint8_t> bytes, size_t depth);
148std::pair<std::string, Attribute> attribute(std::span<const uint8_t> bytes, size_t depth) {
149 Reader r(bytes);
150 Attribute a;
151 std::string name;
152 while (!r.empty()) {
153 auto f = r.next();
154 switch (f.id) {
155 case 1:
156 wire(f, 2);
157 name = f.text();
158 break;
159 case 20:
160 wire(f, 0);
161 a.type = static_cast<int>(f.number);
162 break;
163 case 2: {
164 wire(f, 5);
165 std::memcpy(&a.real, f.bytes.data(), 4);
166 break;
167 }
168 case 3:
169 wire(f, 0);
170 a.integer = std::bit_cast<int64_t>(f.number);
171 break;
172 case 4:
173 wire(f, 2);
174 a.text = f.text();
175 break;
176 case 5:
177 wire(f, 2);
178 a.tensor = tensor(f.bytes).second;
179 break;
180 case 7: reals(f, a.reals); break;
181 case 8: integers(f, a.integers); break;
182 case 9:
183 wire(f, 2);
184 a.strings.push_back(f.text());
185 break;
186 case 6:
187 wire(f, 2);
188 a.graph = std::make_shared<ModelData>(parseGraph(f.bytes, depth + 1));
189 break;
190 case 11: throw Failure("Multiple graph attributes unsupported", DiagnosticCode::Unsupported);
191 case 21: throw Failure("Function attribute references unsupported", DiagnosticCode::Unsupported);
192 default: break;
193 }
194 }
195 if (name.empty() || !a.type) throw Failure("Incomplete attribute", DiagnosticCode::ParseError);
196 return {std::move(name), std::move(a)};
197}
198Node node(std::span<const uint8_t> bytes, size_t depth) {
199 Reader r(bytes);
200 Node n;
201 while (!r.empty()) {
202 auto f = r.next();
203 if (f.id >= 1 && f.id <= 5) wire(f, 2);
204 switch (f.id) {
205 case 1: n.inputs.push_back(f.text()); break;
206 case 2: n.outputs.push_back(f.text()); break;
207 case 3: n.name = f.text(); break;
208 case 4: n.op = f.text(); break;
209 case 5:
210 if (!n.attrs.insert(attribute(f.bytes, depth)).second)
211 throw Failure("Duplicate attribute", DiagnosticCode::ParseError);
212 break;
213 case 7:
214 wire(f, 2);
215 n.domain = f.text();
216 break;
217 case 8: throw Failure("Function overloads unsupported", DiagnosticCode::Unsupported);
218 default: break;
219 }
220 }
221 if (n.op.empty() || n.outputs.empty()) throw Failure("Incomplete node", DiagnosticCode::ParseError);
222 return n;
223}
224std::pair<std::string, Input> input(std::span<const uint8_t> bytes) {
225 Reader r(bytes);
226 std::string name;
227 Input value;
228 while (!r.empty()) {
229 auto f = r.next();
230 if (f.id == 1) {
231 wire(f, 2);
232 name = f.text();
233 }
234 if (f.id == 2) {
235 wire(f, 2);
236 Reader type(f.bytes);
237 while (!type.empty()) {
238 auto t = type.next();
239 if (t.id == 4) {
240 value.type = 0;
241 continue;
242 }
243 if (t.id != 1) throw Failure("Non-tensor graph input", DiagnosticCode::Unsupported);
244 wire(t, 2);
245 Reader tt(t.bytes);
246 while (!tt.empty()) {
247 auto p = tt.next();
248 if (p.id == 1) {
249 wire(p, 0);
250 value.type = static_cast<int>(p.number);
251 }
252 if (p.id == 2) {
253 wire(p, 2);
254 Reader shape(p.bytes);
255 while (!shape.empty()) {
256 auto d = shape.next();
257 if (d.id != 1) continue;
258 wire(d, 2);
259 Reader dim(d.bytes);
260 int64_t v = -1;
261 while (!dim.empty()) {
262 auto x = dim.next();
263 if (x.id == 1) {
264 wire(x, 0);
265 v = std::bit_cast<int64_t>(x.number);
266 }
267 }
268 value.shape.push_back(v);
269 }
270 }
271 }
272 }
273 }
274 }
275 if (name.empty()) throw Failure("Empty value name", DiagnosticCode::ParseError);
276 if (value.shape.size() > 6) throw Failure("Input rank exceeds six", DiagnosticCode::Unsupported);
277 for (auto d : value.shape)
278 if (d < -1 || d > INT32_MAX) throw Failure("Invalid interface dimension", DiagnosticCode::ParseError);
279 return {name, value};
280}
281ModelData parseGraph(std::span<const uint8_t> bytes, size_t depth) {
282 ModelData model;
283 if (depth > 16) throw Failure("ONNX graph nesting limit exceeded", DiagnosticCode::ParseError);
284 Reader g(bytes);
285 while (!g.empty()) {
286 auto f = g.next();
287 if (f.id == 1) {
288 wire(f, 2);
289 if (model.nodes.size() >= 100000) throw Failure("Node limit exceeded", DiagnosticCode::ParseError);
290 model.nodes.push_back(node(f.bytes, depth));
291 }
292 if (f.id == 5) {
293 wire(f, 2);
294 auto t = tensor(f.bytes);
295 t.second.bytes.retainAcrossRuns();
296 if (t.first.empty() || !model.constants.emplace(std::move(t)).second)
297 throw Failure("Duplicate/empty initializer", DiagnosticCode::ParseError);
298 }
299 if (f.id == 11) {
300 wire(f, 2);
301 auto v = input(f.bytes);
302 model.info.inputs.push_back(v.first);
303 if (!model.inputs.emplace(std::move(v)).second)
304 throw Failure("Duplicate graph input", DiagnosticCode::ParseError);
305 }
306 if (f.id == 12) {
307 wire(f, 2);
308 model.info.outputs.push_back(input(f.bytes).first);
309 }
310 if (f.id == 15) throw Failure("Sparse initializers unsupported", DiagnosticCode::Unsupported);
311 }
312 return model;
313}
314} // namespace
315
316ModelData parse(std::span<const uint8_t> bytes) {
317 if constexpr (std::endian::native != std::endian::little)
318 throw Failure("ONNX requires little-endian host", DiagnosticCode::Unsupported);
319 if (bytes.empty() || bytes.size() > 512u * 1024u * 1024u)
320 throw Failure("Model size exceeds import limit", DiagnosticCode::ParseError);
322 Reader r(bytes);
323 std::span<const uint8_t> graph;
324 while (!r.empty()) {
325 auto f = r.next();
326 if (f.id == 1) {
327 wire(f, 0);
328 model.info.irVersion = static_cast<int64_t>(f.number);
329 }
330 if (f.id == 7) {
331 wire(f, 2);
332 if (!graph.empty()) throw Failure("Duplicate graph", DiagnosticCode::ParseError);
333 graph = f.bytes;
334 }
335 if (f.id == 8) {
336 wire(f, 2);
337 Reader ops(f.bytes);
338 std::string domain;
339 int64_t version = 0;
340 while (!ops.empty()) {
341 auto x = ops.next();
342 if (x.id == 1) {
343 wire(x, 2);
344 domain = x.text();
345 }
346 if (x.id == 2) {
347 wire(x, 0);
348 version = x.number;
349 }
350 }
351 if (!model.opsets.emplace(domain, version).second)
352 throw Failure("Duplicate opset", DiagnosticCode::ParseError);
353 }
354 if (f.id == 20 || f.id == 25 || f.id == 26)
355 throw Failure("Training/functions/device configuration unsupported", DiagnosticCode::Unsupported);
356 }
357 if (model.info.irVersion < 3 || model.info.irVersion > 10 || !model.opsets.contains("") ||
358 model.opsets.at("") < 13 || model.opsets.at("") > 17)
359 throw Failure("Unsupported ONNX IR/opset version", DiagnosticCode::UnknownVersion);
360 if (model.opsets.contains("com.microsoft") && model.opsets.at("com.microsoft") != 1)
361 throw Failure("Unsupported Microsoft opset", DiagnosticCode::UnknownVersion);
362 if (graph.empty()) throw Failure("Missing ONNX graph", DiagnosticCode::ParseError);
363 auto parsed = parseGraph(graph, 0);
364 model.nodes = std::move(parsed.nodes);
365 model.constants = std::move(parsed.constants);
366 model.inputs = std::move(parsed.inputs);
367 model.info.inputs = std::move(parsed.info.inputs);
368 model.info.outputs = std::move(parsed.info.outputs);
369 std::set<std::string> names;
370 for (const auto& [name, v] : model.constants) {
371 names.insert(name);
372 model.info.initializerBytes += v.bytes.size();
373 }
374 for (const auto& [name, v] : model.inputs) names.insert(name);
375 for (const auto& n : model.nodes) {
376 if (!model.opsets.contains(n.domain)) throw Failure("Node domain has no opset", DiagnosticCode::ParseError);
377 for (const auto& name : n.inputs)
378 if (!name.empty() && !names.contains(name))
379 throw Failure("Input not in topological order: " + name, DiagnosticCode::ParseError);
380 for (const auto& name : n.outputs)
381 if (!name.empty() && !names.insert(name).second)
382 throw Failure("Duplicate graph value: " + name, DiagnosticCode::ParseError);
383 if (!isSupported(n)) model.info.unsupportedNodes.push_back(n.name + " [" + n.domain + "::" + n.op + "]");
384 }
385 for (const auto& out : model.info.outputs)
386 if (!names.contains(out)) throw Failure("Undefined graph output", DiagnosticCode::ParseError);
387 size_t recursiveNodes = 0;
388 std::function<void(const ModelData&, const std::set<std::string>&, bool)> validateNested;
389 validateNested = [&](const ModelData& g, const std::set<std::string>& outer, bool root) {
390 recursiveNodes += g.nodes.size();
391 if (recursiveNodes > 100000) throw Failure("Recursive node limit exceeded", DiagnosticCode::ParseError);
392 auto available = outer;
393 std::set<std::string> locals;
394 for (const auto& [name, t] : g.constants) locals.insert(name);
395 for (const auto& [name, t] : g.inputs) locals.insert(name);
396 for (const auto& n : g.nodes)
397 for (const auto& out : n.outputs)
398 if (!out.empty() && !locals.insert(out).second)
399 throw Failure("Duplicate nested graph value: " + out, DiagnosticCode::ParseError);
400 available.insert(locals.begin(), locals.end());
401 for (const auto& n : g.nodes) {
402 if (!model.opsets.contains(n.domain))
403 throw Failure("Nested node domain has no opset", DiagnosticCode::ParseError);
404 for (const auto& input : n.inputs)
405 if (!input.empty() && !available.contains(input))
406 throw Failure("Undefined lexical capture: " + input, DiagnosticCode::ParseError);
407 if (!root && !isSupported(n))
408 model.info.unsupportedNodes.push_back(n.name + " [" + n.domain + "::" + n.op + "]");
409 for (const auto& [key, a] : n.attrs)
410 if (a.graph) validateNested(*a.graph, available, false);
411 }
412 for (const auto& out : g.info.outputs)
413 if (!available.contains(out)) throw Failure("Undefined nested output: " + out, DiagnosticCode::ParseError);
414 };
415 validateNested(model, {}, true);
416 model.info.nodeCount = model.nodes.size();
417 return model;
418}
419} // namespace eve::tensor::onnx_detail
double value
float x
Definition AnimClip.cpp:738
int root
Definition AnimSmr.cpp:119
glm::vec4 p[6]
EvpackChunkInput input
Definition Evpack.cpp:170
tensor::Graph g
Definition GpuGraph.cpp:7
std::uint32_t key
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
double r
float v
std::string text
std::uint64_t bytes
std::string name
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
const std::string * tag
const RuntimeTensor & tensor
Definition OnnxLstm.cpp:26
int wire
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
float f
float d
float t
glm::mat4 model
const RoadNode * node
double number
std::uint32_t count
float size
Definition TreeMesh.cpp:156
std::uint32_t depth
constexpr HexDirection next(HexDirection d) noexcept
The next direction clockwise (NW wraps to NE).
Definition HexMetrics.h:76
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.
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