载入中...
搜索中...
未找到
OnnxRuntime.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <functional>
4#include <queue>
5#include <random>
6#include <set>
8
10namespace {
11struct Value {
12 const RuntimeTensor* borrowed = nullptr;
13 std::shared_ptr<const RuntimeTensor> owned;
14 std::shared_ptr<std::vector<Value>> sequence;
16 const RuntimeTensor& tensor() const {
17 if (owned) return *owned;
18 if (borrowed) return *borrowed;
19 throw Failure("Expected tensor, received sequence");
20 }
21};
22using Values = std::unordered_map<std::string, Value>;
23struct Context {
24 OnnxCompute* compute;
25 uint64_t seed;
26 bool finite;
27 size_t steps = 0, liveBytes = 0;
28 std::unordered_map<const Node*, std::mt19937> randomStreams;
29};
30Value hold(RuntimeTensor t, Context& c) {
31 validate(t);
32 if (c.finite && t.element == OnnxElement::Float32)
33 for (size_t i = 0; i < count(t.shape); ++i)
34 if (!std::isfinite(read<float>(t, i))) throw Failure("Nonfinite float output");
35 const auto size = t.bytes.size();
36 if (size > 512u * 1024u * 1024u - c.liveBytes) throw Failure("ONNX live tensor memory limit exceeded");
37 // Context encloses every runtime Value; escaping public outputs are owning copies.
38 auto* raw = new RuntimeTensor(std::move(t));
39 c.liveBytes += size;
40 return {nullptr,
41 std::shared_ptr<const RuntimeTensor>(raw,
42 [&c, size](const RuntimeTensor* p) {
43 delete p;
44 c.liveBytes -= size;
45 }),
46 {}};
47}
48std::set<std::string> captures(const ModelData& graph);
49std::set<std::string> dependencies(const Node& n) {
50 std::set<std::string> result;
51 for (const auto& i : n.inputs)
52 if (!i.empty()) result.insert(i);
53 for (const auto& [key, a] : n.attrs)
54 if (a.graph) {
55 auto names = captures(*a.graph);
56 result.insert(names.begin(), names.end());
57 }
58 return result;
59}
60std::set<std::string> captures(const ModelData& graph) {
61 std::set<std::string> locals, result;
62 for (const auto& [name, t] : graph.constants) locals.insert(name);
63 for (const auto& [name, t] : graph.inputs) locals.insert(name);
64 for (const auto& n : graph.nodes)
65 for (const auto& out : n.outputs)
66 if (!out.empty()) locals.insert(out);
67 for (const auto& n : graph.nodes) {
68 auto deps = dependencies(n);
69 for (const auto& d : deps)
70 if (!locals.contains(d)) result.insert(d);
71 }
72 for (const auto& out : graph.info.outputs)
73 if (!locals.contains(out)) result.insert(out);
74 return result;
75}
76void admit(const Node& n) {
77 if (!isSupported(n))
78 throw Failure(n.name + ": Unsupported ONNX node " + n.domain + "::" + n.op, DiagnosticCode::Unsupported,
79 n.name);
80 for (const auto& [key, a] : n.attrs)
81 if (a.graph)
82 for (const auto& child : a.graph->nodes) admit(child);
83}
84std::vector<size_t> plan(const ModelData& g, const std::vector<std::string>& requested) {
85 std::unordered_map<std::string, size_t> producers;
86 for (size_t i = 0; i < g.nodes.size(); ++i)
87 for (const auto& name : g.nodes[i].outputs)
88 if (!name.empty() && !producers.emplace(name, i).second) throw Failure("Duplicate graph output");
89 std::vector<bool> selected(g.nodes.size(), false);
90 std::vector<std::string> pending = requested;
91 while (!pending.empty()) {
92 auto name = std::move(pending.back());
93 pending.pop_back();
94 auto p = producers.find(name);
95 if (p == producers.end() || selected[p->second]) continue;
96 selected[p->second] = true;
97 admit(g.nodes[p->second]);
98 auto deps = dependencies(g.nodes[p->second]);
99 pending.insert(pending.end(), deps.begin(), deps.end());
100 }
101 std::vector<size_t> indegree(g.nodes.size());
102 std::vector<std::vector<size_t>> users(g.nodes.size());
103 size_t total = 0;
104 std::priority_queue<size_t, std::vector<size_t>, std::greater<size_t>> ready;
105 for (size_t i = 0; i < g.nodes.size(); ++i)
106 if (selected[i]) {
107 ++total;
108 std::set<size_t> parents;
109 for (const auto& name : dependencies(g.nodes[i])) {
110 auto p = producers.find(name);
111 if (p != producers.end()) parents.insert(p->second);
112 }
113 indegree[i] = parents.size();
114 for (auto p : parents) users[p].push_back(i);
115 if (parents.empty()) ready.push(i);
116 }
117 std::vector<size_t> order;
118 while (!ready.empty()) {
119 size_t i = ready.top();
120 ready.pop();
121 order.push_back(i);
122 for (auto user : users[i])
123 if (--indegree[user] == 0) ready.push(user);
124 }
125 if (order.size() != total) throw Failure("ONNX graph has cyclic lexical dependencies", DiagnosticCode::ParseError);
126 return order;
127}
128std::vector<Value> run(const ModelData&, Values, const std::vector<std::string>&, Context&, size_t);
129std::vector<Value> control(const Node& n, const std::vector<Value>& in, const Values& values, Context& c,
130 size_t depth) {
131 auto scalar = [&](size_t i) {
132 const auto& t = in.at(i).tensor();
133 if (count(t.shape) != 1) throw Failure("Control input must contain one element");
134 return integer(t);
135 };
136 auto seq = [&](size_t i) -> const std::vector<Value>& {
137 if (!in.at(i).sequence) throw Failure("Expected sequence input");
138 return *in[i].sequence;
139 };
140 auto sequence = [](std::vector<Value> v, OnnxElement type) {
141 return Value{nullptr, {}, std::make_shared<std::vector<Value>>(std::move(v)), type};
142 };
143 if (n.op == "SequenceEmpty") {
144 auto type = static_cast<OnnxElement>(attr(n, "dtype", 1));
146 return {sequence({}, type)};
147 }
148 if (n.op == "SequenceAt") {
149 const auto& s = seq(0);
150 int64_t i = scalar(1);
151 if (i < 0) i += static_cast<int64_t>(s.size());
152 if (i < 0 || static_cast<size_t>(i) >= s.size()) throw Failure("SequenceAt out of bounds");
153 return {s[i]};
154 }
155 if (n.op == "SequenceInsert") {
156 auto s = seq(0);
157 if (s.size() >= 100000) throw Failure("Sequence length limit exceeded");
158 int64_t i = in.size() > 2 ? scalar(2) : static_cast<int64_t>(s.size());
159 if (i < 0) i += static_cast<int64_t>(s.size());
160 if (i < 0 || static_cast<size_t>(i) > s.size()) throw Failure("SequenceInsert out of bounds");
161 const auto& t = in.at(1).tensor();
162 if (in[0].sequenceElement != t.element) throw Failure("Sequence element dtype mismatch");
163 s.insert(s.begin() + i, in[1]);
164 return {sequence(std::move(s), in[0].sequenceElement)};
165 }
166 if (n.op == "SplitToSequence") {
167 const auto& x = in.at(0).tensor();
168 int a = axis(attr(n, "axis", 0), x.shape.size());
169 std::vector<int64_t> lengths;
170 if (in.size() > 1 && (in[1].borrowed || in[1].owned)) {
171 const auto& split = in[1].tensor();
172 lengths = ints(split);
173 if (split.shape.empty()) {
174 if (lengths[0] <= 0) throw Failure("Invalid split size");
175 int64_t width = lengths[0];
176 lengths.clear();
177 for (int64_t i = 0; i < x.shape[a]; i += width) {
178 if (lengths.size() >= 100000) throw Failure("Sequence length limit exceeded");
179 lengths.push_back(std::min(width, x.shape[a] - i));
180 }
181 }
182 } else
183 lengths.assign(x.shape[a], 1);
184 if (lengths.size() > 100000) throw Failure("Sequence length limit exceeded");
185 int64_t total = 0;
186 for (auto size : lengths) {
187 if (size < 0 || size > x.shape[a] - total) throw Failure("Invalid split lengths");
188 total += size;
189 }
190 if (total != x.shape[a]) throw Failure("Split lengths do not cover axis");
191 size_t outer = 1, inner = elementSize(x.element);
192 for (int j = 0; j < a; ++j) outer *= x.shape[j];
193 for (size_t j = a + 1; j < x.shape.size(); ++j) inner *= x.shape[j];
194 std::vector<Value> s;
195 size_t offset = 0;
196 for (auto length : lengths) {
197 auto shape = x.shape;
198 shape[a] = length;
199 if (!attr(n, "keepdims", 1)) {
200 if (in.size() > 1 && (in[1].borrowed || in[1].owned))
201 throw Failure("keepdims=0 with explicit split unsupported", DiagnosticCode::Unsupported);
202 shape.erase(shape.begin() + a);
203 }
204 RuntimeTensor out{x.element, shape, std::vector<uint8_t>(count(shape) * elementSize(x.element))};
205 for (size_t o = 0; o < outer; ++o)
206 if (length)
207 std::memcpy(out.bytes.data() + o * length * inner,
208 x.bytes.data() + (o * x.shape[a] + offset) * inner, length * inner);
209 s.push_back(hold(std::move(out), c));
210 offset += length;
211 }
212 return {sequence(std::move(s), x.element)};
213 }
214 if (n.op == "ConcatFromSequence") {
215 const auto& s = seq(0);
216 if (s.empty()) throw Failure("Cannot concatenate empty sequence");
217 std::vector<RuntimeTensor> tensors;
218 std::vector<const RuntimeTensor*> inputs;
219 for (const auto& v : s) {
220 tensors.push_back(v.tensor());
221 if (attr(n, "new_axis", 0)) {
222 const int a = axis(attr(n, "axis", 0), tensors.back().shape.size() + 1);
223 tensors.back().shape.insert(tensors.back().shape.begin() + a, 1);
224 }
225 }
226 for (const auto& t : tensors) inputs.push_back(&t);
227 Node concat = n;
228 concat.op = "Concat";
229 auto out = executeShape(concat, inputs);
230 if (!out) throw Failure("ConcatFromSequence dispatch failed");
231 return {hold(std::move(*out), c)};
232 }
233 if (n.op == "If") {
234 if (in.at(0).tensor().element != OnnxElement::Bool) throw Failure("If condition must be boolean");
235 auto key = scalar(0) ? "then_branch" : "else_branch";
236 const auto& g = n.attrs.at(key).graph;
237 if (!g) throw Failure("Missing If branch");
238 return run(*g, values, g->info.outputs, c, depth + 1);
239 }
240 if (n.op == "Loop") {
241 const auto& body = n.attrs.at("body").graph;
242 if (!body) throw Failure("Missing Loop body");
243 if (in.size() < 2 || body->info.inputs.size() != in.size() || body->info.outputs.size() != in.size() - 1 ||
244 n.outputs.size() != in.size() - 2)
245 throw Failure("Loop scan outputs unsupported", DiagnosticCode::Unsupported);
246 const int64_t trips = (in[0].borrowed || in[0].owned) ? scalar(0) : 10000;
247 if (trips < 0 || trips > 10000) throw Failure("Loop trip budget exceeded");
248 bool condition = (in[1].borrowed || in[1].owned) ? scalar(1) != 0 : true;
249 std::vector<Value> state(in.begin() + 2, in.end());
250 for (int64_t i = 0; i < trips && condition; ++i) {
251 auto scope = values;
252 scope[body->info.inputs[0]] = hold(make(OnnxElement::Int64, {}, std::vector<int64_t>{i}), c);
253 scope[body->info.inputs[1]] =
254 hold(make(OnnxElement::Bool, {}, std::vector<uint8_t>{uint8_t(condition)}), c);
255 for (size_t j = 0; j < state.size(); ++j) scope[body->info.inputs[j + 2]] = state[j];
256 auto result = run(*body, std::move(scope), body->info.outputs, c, depth + 1);
257 const auto& cond = result[0].tensor();
258 if (cond.element != OnnxElement::Bool || count(cond.shape) != 1)
259 throw Failure("Loop returned invalid condition");
260 condition = integer(cond) != 0;
261 state.assign(result.begin() + 1, result.end());
262 }
263 return state;
264 }
265 if (n.op == "RandomUniformLike" || n.op == "RandomNormalLike") {
266 const auto& x = in.at(0).tensor();
267 if (attr(n, "dtype", static_cast<int>(x.element)) != 1)
268 throw Failure("Random output requires FP32", DiagnosticCode::Unsupported);
269 uint64_t seed = c.seed;
270 for (unsigned char ch : n.name) seed = (seed ^ ch) * 1099511628211ull;
271 if (n.attrs.contains("seed")) {
272 uint32_t bits;
273 std::memcpy(&bits, &n.attrs.at("seed").real, 4);
274 seed ^= bits;
275 }
276 auto [stream, inserted] = c.randomStreams.try_emplace(&n, static_cast<uint32_t>(seed ^ (seed >> 32)));
277 auto& rng = stream->second;
278 auto uniform = [&]() { return (double(rng()) + .5) / 4294967296.; };
279 std::vector<float> v(count(x.shape));
280 const float low = n.attrs.contains("low") ? n.attrs.at("low").real : 0,
281 high = n.attrs.contains("high") ? n.attrs.at("high").real : 1,
282 mean = n.attrs.contains("mean") ? n.attrs.at("mean").real : 0,
283 scale = n.attrs.contains("scale") ? n.attrs.at("scale").real : 1;
284 if (!std::isfinite(low) || !std::isfinite(high) || !std::isfinite(mean) || !std::isfinite(scale) ||
285 high < low || scale < 0)
286 throw Failure("Invalid random distribution parameters");
287 for (auto& f : v)
288 f = n.op == "RandomUniformLike" ? static_cast<float>(low + (high - low) * uniform())
289 : static_cast<float>(mean + scale * std::sqrt(-2 * std::log(uniform())) *
290 std::cos(6.283185307179586 * uniform()));
291 return {hold(make(OnnxElement::Float32, x.shape, v), c)};
292 }
293 std::vector<const RuntimeTensor*> tensors;
294 for (const auto& v : in) tensors.push_back(v.borrowed || v.owned ? &v.tensor() : nullptr);
295 auto outputs = execute(n, tensors, c.compute);
296 std::vector<Value> result;
297 for (auto& t : outputs) result.push_back(hold(std::move(t), c));
298 return result;
299}
300std::vector<Value> run(const ModelData& g, Values values, const std::vector<std::string>& requested, Context& c,
301 size_t depth) {
302 if (depth > 16) throw Failure("Execution graph nesting limit exceeded");
303 for (const auto& [name, t] : g.constants)
304 if (!g.inputs.contains(name) || !values.contains(name)) values[name] = {&t, {}, {}};
305 auto order = plan(g, requested);
306 std::unordered_map<std::string, size_t> uses;
307 for (auto i : order)
308 for (const auto& dep : dependencies(g.nodes[i])) ++uses[dep];
309 for (const auto& name : requested) ++uses[name];
310 for (auto i : order) {
311 const auto& n = g.nodes[i];
312 if (++c.steps > 1000000) throw Failure("ONNX execution step budget exceeded");
313 try {
314 std::vector<Value> in;
315 for (const auto& name : n.inputs) {
316 if (name.empty()) {
317 in.push_back({});
318 continue;
319 }
320 auto it = values.find(name);
321 if (it == values.end()) throw Failure("Missing feed or lexical capture: " + name);
322 in.push_back(it->second);
323 }
324 auto outputs = control(n, in, values, c, depth);
325 if (outputs.size() < n.outputs.size()) throw Failure("Operator output arity mismatch");
326 for (size_t j = 0; j < n.outputs.size(); ++j)
327 if (!n.outputs[j].empty() && uses[n.outputs[j]]) values[n.outputs[j]] = std::move(outputs[j]);
328 for (const auto& dep : dependencies(n))
329 if (--uses[dep] == 0) values.erase(dep);
330 } catch (const Failure& e) {
331 throw Failure(n.name + ": " + e.what(), e.code, e.path.empty() ? n.name : e.path);
332 }
333 }
334 std::vector<Value> out;
335 for (const auto& name : requested) {
336 auto it = values.find(name);
337 if (it == values.end()) throw Failure("Missing output: " + name, DiagnosticCode::NotFound);
338 out.push_back(it->second);
339 }
340 return out;
341}
342} // namespace
343std::vector<OnnxNamedTensor> evaluate(const ModelData& model, std::span<const OnnxNamedTensor> feeds,
344 const std::vector<std::string>& requested, OnnxCompute* compute,
346 Context c{compute, options.seed, options.requireFinite};
347 Values inputs;
348 for (const auto& f : feeds) inputs[f.name] = hold({f.tensor.element, f.tensor.shape, f.tensor.bytes}, c);
349 auto outputs = run(model, std::move(inputs), requested, c, 0);
350 std::vector<OnnxNamedTensor> result;
351 size_t bytes = c.liveBytes;
352 for (size_t i = 0; i < outputs.size(); ++i) {
353 const auto& t = outputs[i].tensor();
354 if (t.bytes.size() > 512u * 1024u * 1024u - bytes) throw Failure("ONNX output memory limit exceeded");
355 bytes += t.bytes.size();
356 result.push_back({requested[i], {t.element, t.shape, static_cast<std::vector<uint8_t>>(t.bytes)}});
357 }
358 return result;
359}
360} // namespace eve::tensor::onnx_detail
float x
Definition AnimClip.cpp:738
ActiveSource owned
const std::string & s
std::vector< QuestEvent > pending
float length
Definition CaveMesh.cpp:94
bool split
Definition CaveMesh.cpp:123
glm::vec4 p[6]
eve::Value condition
std::map< std::string, Var > values
tensor::Graph g
Definition GpuGraph.cpp:7
std::uint32_t key
ShaderImageInput shape
glm::vec3 n
Definition Grass.cpp:63
int inputs
Definition GridGraph.cpp:23
float v
std::int32_t second
std::int32_t c
std::uint32_t width
JobScope scope
size_t offset
std::array< float, 3 > scale
std::uint64_t bytes
std::string name
MeleePoint3 a
Definition MeleeHit.cpp:40
std::vector< std::int32_t > order
std::vector< BvhNode > nodes
std::weak_ptr< Run > run
Definition OnnxGpgpu.cpp:25
std::unique_ptr< gpgpu::Sequence > sequence
Definition OnnxGpgpu.cpp:43
const RuntimeTensor & tensor
Definition OnnxLstm.cpp:26
bool finite
OnnxCompute * compute
std::unordered_map< const Node *, std::mt19937 > randomStreams
size_t liveBytes
OnnxElement sequenceElement
const RuntimeTensor * borrowed
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
float f
std::string path
Definition PlayHost.cpp:110
std::uint32_t seed
Definition PointSet.cpp:807
float d
int steps
float t
glm::mat4 model
std::uint32_t count
const SquirrelValueOptions & options
float size
Definition TreeMesh.cpp:156
std::string body
std::uint32_t depth
GPU execution boundary for native ONNX; retains no model and retains compiled resources for the lifet...
Definition OnnxCompute.h:25
std::variant< std::monostate, std::int64_t, double, std::string, bool > Value
Definition Database.h:26
void concat(const float *const *ins, const int *const *inDims, const int *inRanks, int n, int axis, float *out, const int *outDims, int outRank)
Concat.
int64_t integer(const RuntimeTensor &v, size_t i=0)
Integer.
std::vector< int64_t > ints(const RuntimeTensor &v)
Ints.
int64_t attr(const Node &n, const char *key, int64_t fallback)
Attr.
std::vector< int64_t > attrs(const Node &n, const char *key, std::vector< int64_t > fallback)
Attrs.
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< 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.
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.
OnnxElement
ONNX wire element types; distinct from block-quantized Tensor storage.
Definition OnnxModel.h:18
WidgetDesc child(std::string id, std::vector< WidgetDesc > children, float width, float height)
Scrollable child region with an explicit size.
Definition Widget.cpp:635
Per-call deterministic RNG and optional strict finite-output diagnostic.
Definition OnnxModel.h:37
glm::uvec4 info