26 #include "onnx_model.hpp"
32 static bool integer_tensor(
const onnx::TensorProto& tensor, std::vector<int64_t>& values) {
33 if (tensor.data_type() != onnx::TensorProto::INT64 &&
34 tensor.data_type() != onnx::TensorProto::INT32)
return false;
36 if (tensor.has_raw_data()) {
37 const auto& raw = tensor.raw_data();
38 size_t width = tensor.data_type() == onnx::TensorProto::INT64 ? 8 : 4;
39 casadi_assert(raw.size() % width == 0,
"Invalid integer tensor raw_data size");
40 for (
size_t i = 0; i < raw.size(); i += width) {
42 for (
size_t j = 0; j < width; ++j) {
43 bits |=
static_cast<uint64_t
>(
static_cast<unsigned char>(raw[i + j])) << (8 * j);
47 values.push_back(
static_cast<int64_t
>(bits) -
static_cast<int64_t
>((bits >> 31) << 32));
49 values.push_back(bits >> 63 ? -1 -
static_cast<int64_t
>(~bits) :
50 static_cast<int64_t
>(bits));
53 }
else if (tensor.data_type() == onnx::TensorProto::INT64) {
54 values.assign(tensor.int64_data().begin(), tensor.int64_data().end());
56 values.assign(tensor.int32_data().begin(), tensor.int32_data().end());
62 void Onnx::process_graph_initializers(
63 const onnx::GraphProto& graph,
64 std::map<std::string, MX>& value_map,
67 for (
int i = 0; i < graph.initializer_size(); ++i) {
68 const onnx::TensorProto& tensor = graph.initializer(i);
69 std::string tensor_name = tensor.name();
72 uout() <<
" Processing initializer: " << tensor_name << std::endl;
75 value_map[tensor_name] = MX(tensor_to_dm(tensor));
80 void Onnx::process_graph_inputs(
81 const onnx::GraphProto& graph,
82 std::map<std::string, MX>& value_map,
83 std::vector<MX>& func_inputs,
84 std::vector<std::string>& input_names,
87 for (
int i = 0; i < graph.input_size(); ++i) {
88 const onnx::ValueInfoProto& input = graph.input(i);
89 std::string input_name = input.name();
92 if (value_map.count(input_name)) {
94 uout() <<
" Skipping input '" << input_name
95 <<
"' (it's an initializer)" << std::endl;
101 const onnx::TensorShapeProto& shape =
102 input.type().tensor_type().shape();
103 casadi_int rows = get_dimension(shape, 0);
104 casadi_int cols = get_dimension(shape, 1);
107 uout() <<
" Creating input: " << input_name
108 <<
" [" << rows <<
", " << cols <<
"]" << std::endl;
116 Sparsity ov = input_pattern(graph, input_name);
117 MX mx_input = ov.is_null() ?
MX::sym(input_name, cols, rows) : MX::sym(input_name, ov);
118 value_map[input_name] = mx_input;
119 func_inputs.push_back(mx_input);
120 input_names.push_back(input_name);
125 void Onnx::process_graph_nodes(
126 const onnx::GraphProto& graph,
127 std::map<std::string, MX>& value_map,
130 IntegerConstants integer_constants;
131 for (
const auto& tensor : graph.initializer()) {
132 std::vector<int64_t> values;
133 if (
integer_tensor(tensor, values)) integer_constants[tensor.name()] = std::move(values);
137 std::map<std::string, MX> kron_operand_a_, kron_operand_b_;
139 for (
int i = 0; i < graph.node_size(); ++i) {
140 const onnx::NodeProto& node = graph.node(i);
141 std::string op_type = node.op_type();
144 uout() <<
" Processing node " << i <<
": " << op_type << std::endl;
148 if (op_type ==
"Constant" && node.output_size() == 1) {
149 for (
const auto& attr : node.attribute()) {
150 std::vector<int64_t> values;
151 if (attr.name() ==
"value_int") {
153 }
else if (attr.name() ==
"value_ints") {
154 values.assign(attr.ints().begin(), attr.ints().end());
155 }
else if (attr.name() !=
"value" || !
integer_tensor(attr.t(), values)) {
158 integer_constants[node.output(0)] = std::move(values);
172 if (node.name().rfind(
"kron", 0) == 0) {
173 const std::string& nm = node.name();
174 std::size_t us = nm.rfind(
'_');
175 std::string tag = nm.substr(0, us), role = nm.substr(us + 1);
177 kron_operand_a_[tag] = value_map.at(node.input(0));
181 kron_operand_b_[tag] = value_map.at(node.input(0));
184 if (role ==
"P")
continue;
187 casadi_assert(kron_operand_a_.count(tag) && kron_operand_b_.count(tag),
188 "ONNX import: incomplete kron node group '" + tag +
"'");
189 value_map[node.output(0)] =
MX::kron(kron_operand_a_.at(tag), kron_operand_b_.at(tag));
195 std::vector<MX> node_inputs;
196 for (
int j = 0; j < node.input_size(); ++j) {
197 std::string input_name = node.input(j);
200 if (input_name.empty()) {
201 node_inputs.push_back(MX());
205 casadi_assert(value_map.count(input_name),
206 "Unknown input tensor '" + input_name +
207 "' required by node " + std::to_string(i) +
208 " (op_type: " + op_type +
")");
209 node_inputs.push_back(value_map[input_name]);
216 if (op_type ==
"Split") {
217 casadi_assert(node_inputs.size() >= 1,
"Split requires 1 input");
219 casadi_assert(axis == 0 || axis == 1,
"Split: only axis 0 and 1 supported");
222 std::vector<casadi_int> split_sizes;
223 if (node_inputs.size() >= 2) {
224 DM split_dm =
DM(node_inputs[1]);
225 for (casadi_int k = 0; k < split_dm.nnz(); ++k) {
226 split_sizes.push_back(
static_cast<casadi_int
>(
static_cast<double>(split_dm.nz(k))));
229 for (
int a = 0; a < node.attribute_size(); ++a) {
230 if (node.attribute(a).name() ==
"split") {
231 for (
int k = 0; k < node.attribute(a).ints_size(); ++k) {
232 split_sizes.push_back(node.attribute(a).ints(k));
238 if (split_sizes.empty()) {
240 casadi_int total = (axis == 0) ? node_inputs[0].size2() : node_inputs[0].size1();
241 split_sizes.assign(node.output_size(), total / node.output_size());
245 std::vector<casadi_int> offset = {0};
246 for (casadi_int sz : split_sizes) offset.push_back(offset.back() + sz);
247 std::vector<MX> outputs = (axis == 0) ? horzsplit(node_inputs[0], offset)
248 : vertsplit(node_inputs[0], offset);
250 for (casadi_int j = 0; j < outputs.size(); ++j) value_map[node.output(j)] = outputs[j];
254 }
else if (op_type ==
"Scan") {
256 casadi_assert(body !=
nullptr,
"Scan node requires a 'body' subgraph");
257 casadi_int num_scan_inputs =
get_int_attribute(node,
"num_scan_inputs", node.input_size());
258 casadi_int M = node.input_size() - num_scan_inputs;
261 std::set<std::string> body_in_names, defined;
262 for (
int b = 0; b < body->input_size(); ++b) {
263 body_in_names.insert(body->input(b).name());
264 defined.insert(body->input(b).name());
266 for (
int nn = 0; nn < body->node_size(); ++nn)
267 for (
int oo = 0; oo < body->node(nn).output_size(); ++oo)
268 defined.insert(body->node(nn).output(oo));
269 bool has_capture =
false;
270 for (
int nn = 0; nn < body->node_size() && !has_capture; ++nn)
271 for (
int ii = 0; ii < body->node(nn).input_size(); ++ii) {
272 const std::string& in = body->node(nn).input(ii);
273 if (!in.empty() && !defined.count(in)) { has_capture =
true;
break; }
276 if (M == 0 && !has_capture) {
279 Function base = function_from_graph(*body, op_type +
"_body");
280 casadi_int n = node_inputs[0].size2() / base.size2_in(0);
281 std::vector<MX> outputs =
282 base.map(n)(std::vector<MX>(node_inputs.begin(), node_inputs.end()));
283 for (
int j = 0; j < node.output_size(); ++j) value_map[node.output(j)] = outputs[j];
289 const onnx::NodeProto* call =
nullptr;
290 for (
int nn = 0; nn < body->node_size(); ++nn)
291 if (!body->node(nn).domain().empty()) { call = &body->node(nn);
break; }
292 casadi_assert(call,
"reduce-map Scan body must contain a base function call");
293 casadi_int nin = call->input_size(), nout = call->output_size();
296 std::vector<bool> reduce_in(nin), reduce_out(nout,
false);
297 for (casadi_int j = 0; j < nin; ++j) reduce_in[j] = !body_in_names.count(call->input(j));
298 for (
int nn = 0; nn < body->node_size(); ++nn) {
299 if (body->node(nn).op_type() !=
"Add")
continue;
300 for (
int ii = 0; ii < body->node(nn).input_size(); ++ii)
301 for (casadi_int j = 0; j < nout; ++j)
302 if (body->node(nn).input(ii) == call->output(j)) reduce_out[j] =
true;
306 std::vector<std::pair<casadi_int, casadi_int>> in_shapes(nin);
307 for (casadi_int j = 0; j < nin; ++j) {
309 MX cap = value_map.at(call->input(j));
310 in_shapes[j] = {cap.size1(), cap.size2()};
312 for (
int b = 0; b < body->input_size(); ++b)
313 if (body->input(b).name() == call->input(j)) {
314 const onnx::TensorShapeProto& sh = body->input(b).type().tensor_type().shape();
315 in_shapes[j] = {get_dimension(sh, 0), get_dimension(sh, 1)};
321 const onnx::FunctionProto* fp = find_function(call->op_type(), call->domain());
322 casadi_assert(fp,
"reduce-map base function '" + call->op_type() +
"' not found");
323 Function base = function_from_function_proto(*fp, in_shapes, op_type +
"_base");
326 std::vector<MX> args(nin);
327 casadi_int n = 0, r = 0;
328 for (casadi_int j = 0; j < nin; ++j) {
330 args[j] = value_map.at(call->input(j));
332 args[j] = node_inputs[M + r];
333 if (n == 0) n = args[j].size2() / base.size2_in(j);
337 casadi_assert(n > 0,
"reduce-map import: could not infer map size");
338 std::vector<MX> outs = base.map(n, reduce_in, reduce_out)(args);
341 casadi_int si = 0, ci = 0;
342 for (casadi_int j = 0; j < nout; ++j) {
343 if (reduce_out[j]) value_map[node.output(si++)] = outs[j];
345 value_map[node.output(M + ci++)] = outs[j];
350 }
else if (op_type ==
"If") {
353 casadi_assert(then_b && else_b,
"If node requires then_branch and else_branch");
354 std::vector<MX> t = eval_captured_subgraph(*then_b, value_map);
355 std::vector<MX> e = eval_captured_subgraph(*else_b, value_map);
356 MX cond = node_inputs[0];
357 for (
int j = 0; j < node.output_size(); ++j) {
358 value_map[node.output(j)] =
if_else(cond, t[j], e[j]);
363 }
else if (op_type ==
"Loop") {
364 casadi_error(
"ONNX import: 'Loop' control flow operator is not supported.");
369 std::string node_domain = node.domain();
370 if (!node_domain.empty()) {
371 const onnx::FunctionProto* func_proto = find_function(op_type, node_domain);
373 if (func_proto !=
nullptr) {
375 uout() <<
" Function call to: " << node_domain <<
"." << op_type << std::endl;
380 onnx::GraphProto func_graph;
381 func_graph.set_name(func_proto->name());
382 for (
int n = 0; n < func_proto->node_size(); ++n) {
383 *func_graph.add_node() = func_proto->node(n);
387 std::map<std::string, MX> func_value_map;
388 for (
size_t n = 0; n < node_inputs.size() && n < func_proto->input_size(); ++n) {
389 func_value_map[func_proto->input(n)] = node_inputs[n];
392 process_graph_nodes(func_graph, func_value_map,
false);
395 for (
int n = 0; n < node.output_size() && n < func_proto->output_size(); ++n) {
396 std::string func_output_name = func_proto->output(n);
397 casadi_assert(func_value_map.count(func_output_name),
398 "Function output '" + func_output_name +
"' not found");
399 value_map[node.output(n)] = func_value_map[func_output_name];
402 uout() <<
" -> " << node.output(n) << std::endl;
413 output = process_node_operation(op_type, node, node_inputs, integer_constants);
417 casadi_assert(node.output_size() >= 1,
418 "Node must have at least one output");
420 std::string output_name = node.output(0);
421 value_map[output_name] = output;
424 uout() <<
" -> " << output_name << std::endl;
430 void Onnx::collect_graph_outputs(
431 const onnx::GraphProto& graph,
432 const std::map<std::string, MX>& value_map,
433 std::vector<MX>& func_outputs,
434 std::vector<std::string>& output_names,
435 bool verbose)
const {
437 for (
int i = 0; i < graph.output_size(); ++i) {
438 const onnx::ValueInfoProto& output = graph.output(i);
439 std::string output_name = output.name();
441 casadi_assert(value_map.count(output_name),
442 "Unknown output tensor: " + output_name +
443 ". This usually means the ONNX graph contains unsupported operations.");
447 func_outputs.push_back(value_map.at(output_name));
448 output_names.push_back(output_name);
451 uout() <<
" Graph output: " << output_name << std::endl;
456 const onnx::FunctionProto* Onnx::find_function(
const std::string& name,
457 const std::string& domain)
const {
458 for (
int f = 0; f <
model_.functions_size(); ++f)
459 if (
model_.functions(f).name() == name &&
model_.functions(f).domain() == domain)
460 return &
model_.functions(f);
465 casadi_assert(
has_model_,
"No ONNX model loaded. Call load() first.");
466 return function_from_graph(
model_.graph(), name);
469 Function Onnx::function_from_graph(
const onnx::GraphProto& graph,
470 const std::string& name) {
471 std::map<std::string, MX> value_map;
472 std::vector<MX> inputs, outputs;
473 std::vector<std::string> input_names, output_names;
476 uout() <<
"Building CasADi Function '" << name <<
"' from graph '" << graph.name()
477 <<
"' (" << graph.initializer_size() <<
" initializers, "
478 << graph.input_size() <<
" inputs, " << graph.output_size() <<
" outputs, "
479 << graph.node_size() <<
" nodes)" << std::endl;
482 process_graph_initializers(graph, value_map,
verbose_);
483 process_graph_inputs(graph, value_map, inputs, input_names,
verbose_);
484 process_graph_nodes(graph, value_map,
verbose_);
485 collect_graph_outputs(graph, value_map, outputs, output_names,
verbose_);
487 return Function(name, inputs, outputs, input_names, output_names);
490 Function Onnx::function_from_function_proto(
491 const onnx::FunctionProto& fp,
492 const std::vector<std::pair<casadi_int, casadi_int>>& in_shapes,
493 const std::string& name) {
496 g.set_name(fp.name());
497 for (
int j = 0; j < fp.input_size(); ++j) {
498 onnx::ValueInfoProto* vi = g.add_input();
499 vi->set_name(fp.input(j));
500 onnx::TypeProto::Tensor* tt = vi->mutable_type()->mutable_tensor_type();
501 tt->set_elem_type(real_type());
502 tt->mutable_shape()->add_dim()->set_dim_value(in_shapes[j].first);
503 tt->mutable_shape()->add_dim()->set_dim_value(in_shapes[j].second);
505 for (
int k = 0; k < fp.node_size(); ++k) *g.add_node() = fp.node(k);
506 for (
int j = 0; j < fp.output_size(); ++j) g.add_output()->set_name(fp.output(j));
507 return function_from_graph(g, name);
510 std::vector<MX> Onnx::eval_captured_subgraph(
const onnx::GraphProto& graph,
511 std::map<std::string, MX> scope) {
513 process_graph_initializers(graph, scope,
verbose_);
514 process_graph_nodes(graph, scope,
verbose_);
515 std::vector<MX> outputs;
516 for (
int i = 0; i < graph.output_size(); ++i)
517 outputs.push_back(scope.at(graph.output(i).name()));
524 casadi_assert(m.
is_constant(),
"Expected a constant integer tensor");
525 DM dm =
static_cast<DM>(m);
526 std::vector<casadi_int> v;
527 for (casadi_int k = 0; k < dm.
numel(); ++k)
528 v.push_back(
static_cast<casadi_int
>(dm(k).scalar()));
532 MX Onnx::process_node_operation(
533 const std::string& op_type,
534 const onnx::NodeProto& node,
535 const std::vector<MX>& node_inputs,
536 const IntegerConstants& integer_constants) {
538 auto input_ints = [&](
size_t i) -> std::vector<casadi_int> {
539 auto it = integer_constants.find(node.input(i));
540 if (it == integer_constants.end())
return constant_ints(node_inputs[i]);
541 std::vector<casadi_int> values;
542 for (int64_t v : it->second) {
543 casadi_assert(v >= std::numeric_limits<casadi_int>::min() &&
544 v <= std::numeric_limits<casadi_int>::max(),
"Integer constant out of range: " +
str(v));
545 values.push_back(
static_cast<casadi_int
>(v));
555 const MX& x = node_inputs[0];
556 if (mapping->arity == 1) {
557 casadi_assert(node_inputs.size() >= 1, op_type +
" requires 1 input");
558 switch (mapping->casadi_op) {
559 case OP_SIN:
return sin(x);
560 case OP_COS:
return cos(x);
561 case OP_TAN:
return tan(x);
571 case OP_EXP:
return exp(x);
572 case OP_LOG:
return log(x);
579 case OP_ERF:
return erf(x);
580 case OP_INV:
return 1.0 / x;
589 }
else if (mapping->arity == 2) {
590 casadi_assert(node_inputs.size() >= 2, op_type +
" requires 2 inputs");
591 const MX& y = node_inputs[1];
592 switch (mapping->casadi_op) {
593 case OP_ADD:
return x + y;
594 case OP_SUB:
return x - y;
595 case OP_MUL:
return x * y;
596 case OP_DIV:
return x / y;
597 case OP_POW:
return pow(x, y);
605 if (op_type ==
"Less") {
606 casadi_assert(node_inputs.size() >= 2,
"Less requires 2 inputs");
607 output =
if_else(node_inputs[0] < node_inputs[1], MX(1.0), MX(0.0));
609 }
else if (op_type ==
"Equal") {
610 casadi_assert(node_inputs.size() >= 2,
"Equal requires 2 inputs");
611 output = !ne(node_inputs[0], node_inputs[1]);
613 }
else if (op_type ==
"LessOrEqual") {
614 casadi_assert(node_inputs.size() >= 2,
"LessOrEqual requires 2 inputs");
615 output =
if_else(node_inputs[0] <= node_inputs[1], MX(1.0), MX(0.0));
617 }
else if (op_type ==
"Min") {
618 casadi_assert(node_inputs.size() >= 2,
"Min requires 2 inputs");
619 output = fmin(node_inputs[0], node_inputs[1]);
621 }
else if (op_type ==
"Max") {
622 casadi_assert(node_inputs.size() >= 2,
"Max requires 2 inputs");
623 output = fmax(node_inputs[0], node_inputs[1]);
625 }
else if (op_type ==
"Mod") {
626 casadi_assert(node_inputs.size() >= 2,
"Mod requires 2 inputs");
627 output = fmod(node_inputs[0], node_inputs[1]);
629 }
else if (op_type ==
"ReduceSum") {
631 casadi_assert(node_inputs.size() >= 1,
"ReduceSum requires 1 input");
632 output = sum1(sum2(node_inputs[0]));
634 }
else if (op_type ==
"Not") {
635 casadi_assert(node_inputs.size() >= 1,
"Not requires 1 input");
636 output = logic_not(node_inputs[0]);
638 }
else if (op_type ==
"And") {
639 casadi_assert(node_inputs.size() >= 2,
"And requires 2 inputs");
640 output = logic_and(node_inputs[0], node_inputs[1]);
642 }
else if (op_type ==
"Or") {
643 casadi_assert(node_inputs.size() >= 2,
"Or requires 2 inputs");
644 output = logic_or(node_inputs[0], node_inputs[1]);
646 }
else if (op_type ==
"Where") {
647 casadi_assert(node_inputs.size() >= 3,
"Where requires 3 inputs");
648 output =
if_else(node_inputs[0], node_inputs[1], node_inputs[2]);
650 }
else if (op_type ==
"Identity") {
651 casadi_assert(node_inputs.size() >= 1,
"Identity requires 1 input");
652 output = node_inputs[0];
654 }
else if (op_type ==
"Cast") {
657 casadi_assert(node_inputs.size() >= 1,
"Cast requires 1 input");
658 output = node_inputs[0];
660 }
else if (op_type ==
"MatMul") {
662 casadi_assert(node_inputs.size() >= 2,
"MatMul requires 2 inputs");
663 output = mtimes(node_inputs[1], node_inputs[0]);
665 }
else if (op_type ==
"Gemm") {
668 casadi_assert(node_inputs.size() >= 2,
"Gemm requires at least 2 inputs");
669 MX A =
get_int_attribute(node,
"transB", 0) ? node_inputs[1].T() : node_inputs[1];
670 MX B =
get_int_attribute(node,
"transA", 0) ? node_inputs[0].T() : node_inputs[0];
672 if (node_inputs.size() >= 3) {
676 }
else if (op_type ==
"Sum") {
678 casadi_assert(node_inputs.size() >= 1,
"Sum requires at least 1 input");
679 output = node_inputs[0];
680 for (casadi_int idx = 1; idx < node_inputs.size(); ++idx) output = output + node_inputs[idx];
682 }
else if (op_type ==
"Pad") {
686 casadi_assert(node_inputs.size() >= 2,
"Pad requires data and pads");
687 MX block = node_inputs[0];
688 std::vector<casadi_int> pads =
constant_ints(node_inputs[1]);
689 casadi_int col_off = pads[0], row_off = pads[1];
690 casadi_int br = block.size1(), bc = block.size2();
691 casadi_int
C = col_off + bc + pads[2];
696 if (col_off > 0 || pads[2] > 0) {
697 padded = horzcat(MX(Sparsity(br, col_off)), padded, MX(Sparsity(br, pads[2])));
699 if (row_off > 0 || pads[3] > 0) {
700 padded = vertcat(MX(Sparsity(row_off, C)), padded, MX(Sparsity(pads[3], C)));
704 }
else if (op_type ==
"Einsum") {
707 casadi_assert(node_inputs.size() >= 2,
"Einsum requires 2 inputs");
710 size_t comma = eq.find(
','), arrow = eq.find(
"->");
711 casadi_assert(comma != std::string::npos && arrow != std::string::npos,
712 "ONNX import: only binary Einsum 'a,b->c' is supported");
713 std::string sa = eq.substr(0, comma);
714 std::string sb = eq.substr(comma + 1, arrow - comma - 1);
715 std::string sc = eq.substr(arrow + 2);
721 std::reverse(sa.begin(), sa.end());
722 std::reverse(sb.begin(), sb.end());
723 std::reverse(sc.begin(), sc.end());
729 std::map<char, casadi_int> lsize;
730 for (
int t = 0; t < 2; ++t) {
731 const std::string& s = (t == 0) ? sa : sb;
732 const MX& m = node_inputs[t];
734 lsize[s[0]] = m.numel();
735 }
else if (s.size() >= 2) {
736 lsize[s[0]] = m.size1();
737 lsize[s[1]] = m.size2();
741 std::map<char, casadi_int> lab;
742 casadi_int next_label = -1;
743 for (
char ch : sa + sb + sc)
if (!lab.count(ch)) lab[ch] = next_label--;
746 std::vector<casadi_int> da, db, dc, La, Lb, Lc;
747 for (
char ch : sa) { da.push_back(lsize[ch]); La.push_back(lab[ch]); }
748 for (
char ch : sb) { db.push_back(lsize[ch]); Lb.push_back(lab[ch]); }
749 for (
char ch : sc) { dc.push_back(lsize[ch]); Lc.push_back(lab[ch]); }
750 output = einstein(vec(node_inputs[0]), vec(node_inputs[1]), da, db, dc, La, Lb, Lc);
752 }
else if (op_type ==
"Det") {
753 casadi_assert(node_inputs.size() >= 1,
"Det requires 1 input");
754 output = det(node_inputs[0]);
756 }
else if (op_type ==
"ReduceLogSumExp") {
758 casadi_assert(node_inputs.size() >= 1,
"ReduceLogSumExp requires 1 input");
759 output = log(sum1(sum2(exp(node_inputs[0]))));
761 }
else if (op_type ==
"Constant") {
763 const onnx::AttributeProto* value_attr =
nullptr;
764 for (
int a = 0; a < node.attribute_size(); ++a) {
765 const auto& attr = node.attribute(a);
766 const std::string& an = attr.name();
767 if (an ==
"value" || an ==
"sparse_value" || an ==
"value_int" ||
768 an ==
"value_ints" || an ==
"value_float" || an ==
"value_floats") {
769 casadi_assert(value_attr ==
nullptr,
"Constant node has multiple value attributes");
773 casadi_assert(value_attr !=
nullptr,
774 "Constant node requires a numeric value or sparse_value attribute");
775 const auto& attr = *value_attr;
776 if (attr.name() ==
"sparse_value") {
777 output = MX(sparse_tensor_to_dm(attr.sparse_tensor()));
778 }
else if (attr.name() ==
"value") {
779 output = MX(tensor_to_dm(attr.t()));
780 }
else if (attr.name() ==
"value_int") {
781 output = MX(
static_cast<double>(attr.i()));
782 }
else if (attr.name() ==
"value_float") {
783 output = MX(
static_cast<double>(attr.f()));
785 std::vector<double> values;
786 if (attr.name() ==
"value_ints") {
787 for (
auto v : attr.ints()) values.push_back(
static_cast<double>(v));
789 for (
auto v : attr.floats()) values.push_back(
static_cast<double>(v));
791 output = MX(
DM(values));
795 }
else if (op_type ==
"Reshape") {
796 casadi_assert(node_inputs.size() >= 2,
797 "Reshape operation requires 2 inputs (data and shape)");
799 casadi_assert(node_inputs[1].is_constant(),
800 "Reshape shape must be a constant");
801 std::vector<casadi_int> shape = input_ints(1);
807 if (shape.size() > 2) {
809 output = densify(node_inputs[0]);
812 casadi_int s0 = shape.empty() ? 1 : shape[0];
813 casadi_int s1 = shape.size() > 1 ? shape[1] : 1;
814 output = reshape(densify(node_inputs[0]), s1, s0);
817 }
else if (op_type ==
"Unsqueeze") {
818 casadi_assert(!node_inputs.empty() && node_inputs[0].is_vector(),
819 "Unsqueeze: only scalar and vector operands are supported");
820 std::vector<casadi_int> axes;
821 if (node_inputs.size() > 1) {
824 for (
const auto& attr : node.attribute()) {
825 if (attr.name() ==
"axes") axes.assign(attr.ints().begin(), attr.ints().end());
828 casadi_assert(axes.size() == 1 && axes[0] >= -2 && axes[0] < 2,
829 "Unsqueeze: expected one axis for a 2-D result");
830 casadi_int axis = axes[0] < 0 ? axes[0] + 2 : axes[0];
831 output = axis == 0 ? reshape(node_inputs[0], node_inputs[0].numel(), 1) :
832 reshape(node_inputs[0], 1, node_inputs[0].numel());
834 }
else if (op_type ==
"Concat") {
838 output = horzcat(node_inputs);
839 }
else if (axis == 1) {
840 output = vertcat(node_inputs);
842 casadi_error(
"Concat with axis=" + std::to_string(axis) +
843 " not supported. Only axis=0 (vertcat) and axis=1 (horzcat) are supported.");
846 }
else if (op_type ==
"Slice") {
848 casadi_assert(node_inputs.size() >= 3,
849 "Slice requires at least 3 inputs (data, starts, ends)");
851 MX data = node_inputs[0];
852 casadi_assert(node_inputs[1].is_constant() && node_inputs[2].is_constant(),
853 "Slice starts and ends must be constants");
855 DM starts_dm =
static_cast<DM>(node_inputs[1]);
856 DM ends_dm =
static_cast<DM>(node_inputs[2]);
859 std::vector<casadi_int> axes, steps;
860 if (node_inputs.size() >= 4 && !node_inputs[3].is_empty()) {
861 axes = input_ints(3);
863 for (casadi_int k = 0; k < starts_dm.numel(); ++k) axes.push_back(k);
865 if (node_inputs.size() >= 5 && !node_inputs[4].is_empty()) {
866 steps = input_ints(4);
868 steps.assign(starts_dm.numel(), 1);
871 casadi_assert(axes.size() <= 2,
"Slice: only up to 2D slicing supported");
872 casadi_assert(starts_dm.numel() == axes.size() && ends_dm.numel() == axes.size() &&
873 steps.size() == axes.size(),
"Slice: starts, ends, axes and steps must have equal lengths");
876 Slice row_slice, col_slice;
877 casadi_int nrow = data.size1(), ncol = data.size2();
878 for (casadi_int a = 0; a < static_cast<casadi_int>(axes.size()); ++a) {
879 casadi_int axis = axes[a] < 0 ? axes[a] + 2 : axes[a];
880 casadi_assert(axis == 0 || axis == 1,
"Slice: axis out of range");
881 casadi_int n = axis == 0 ? ncol : nrow;
882 casadi_int sp = steps[a];
883 casadi_assert(sp != 0,
"Slice: step must not be zero");
884 sp = std::max(-std::max(n, casadi_int(1)), std::min(sp, std::max(n, casadi_int(1))));
885 auto bound = [&](
size_t input,
const DM& dm) -> casadi_int {
886 auto it = integer_constants.find(node.input(input));
887 if (it != integer_constants.end()) {
888 int64_t v = it->second.at(a);
890 return static_cast<casadi_int
>(std::max<int64_t>(sp > 0 ? 0 : -1,
891 std::min<int64_t>(v, sp > 0 ? n : n - 1)));
894 double v = dm(a).scalar();
896 return static_cast<casadi_int
>(std::max(sp > 0 ? 0. : -1.,
897 std::min(v,
static_cast<double>(sp > 0 ? n : n - 1))));
899 casadi_int st = bound(1, starts_dm);
900 casadi_int en = bound(2, ends_dm);
902 Slice indices = (sp > 0 ? st >= en : st <= en) ? Slice(0, 0) :
903 Slice(st, en == -1 ? std::numeric_limits<casadi_int>::max() : en, sp);
910 output = data(row_slice, col_slice);
912 }
else if (op_type ==
"Gather") {
914 casadi_assert(node_inputs.size() >= 2,
"Gather requires data and indices");
917 MX data = densify(node_inputs[0]);
918 MX indices_mx = node_inputs[1];
921 casadi_assert(indices_mx.is_constant(),
"Gather indices must be constant");
922 std::vector<casadi_int> indices = input_ints(1);
925 if (data.size2() == 1 || axis == 1) {
926 output = data(indices, Slice());
927 }
else if (axis == 0) {
928 output = data(Slice(), indices);
930 casadi_error(
"Gather: only axis 0 and 1 supported for 2D tensors");
933 }
else if (op_type ==
"ScatterElements") {
936 casadi_assert(node_inputs.size() >= 3,
"ScatterElements requires data, indices, updates");
937 MX se_data = densify(node_inputs[0]);
938 MX se_idx = node_inputs[1];
939 MX se_upd = vec(node_inputs[2]);
945 MX w =
MX::sym(
"scel_w", se_data.numel(), 1), g;
946 if (se_idx.is_constant()) {
949 w.get_nz(g,
false, vec(se_idx));
951 output = se_data + reshape(jtimes(g, w, se_upd,
true), se_data.size1(), se_data.size2());
954 if (se_idx.is_constant()) {
957 output.set_nz(se_upd,
false, vec(se_idx));
961 }
else if (op_type ==
"ScatterND") {
965 casadi_assert(node_inputs.size() >= 3,
"ScatterND requires data, indices, updates");
966 MX data = densify(node_inputs[0]);
967 std::vector<casadi_int> coords =
constant_ints(node_inputs[1]);
968 casadi_int R = data.size1();
969 std::vector<casadi_int> pos(coords.size() / 2);
970 for (casadi_int i = 0; i < static_cast<casadi_int>(pos.size()); ++i) {
971 pos[i] = coords[2 * i] * R + coords[2 * i + 1];
978 MX w =
MX::sym(
"addnz_w", data.numel(), 1), g;
980 MX scatter = jtimes(g, w, vec(node_inputs[2]),
true);
981 output = data + reshape(scatter, data.size1(), data.size2());
987 }
else if (op_type ==
"Tile") {
989 casadi_assert(node_inputs.size() >= 2,
"Tile requires data and repeats inputs");
990 MX data = node_inputs[0];
991 MX repeats_mx = node_inputs[1];
992 casadi_assert(repeats_mx.is_constant(),
"Tile repeats must be constant");
993 DM repeats_dm =
static_cast<DM>(repeats_mx);
995 casadi_int rows_repeat = 1, cols_repeat = 1;
996 if (repeats_dm.numel() >= 1) rows_repeat =
static_cast<casadi_int
>(repeats_dm(0).scalar());
997 if (repeats_dm.numel() >= 2) cols_repeat =
static_cast<casadi_int
>(repeats_dm(1).scalar());
1000 output = repmat(data, cols_repeat, rows_repeat);
1002 }
else if (op_type ==
"GatherElements") {
1007 casadi_assert(node_inputs.size() >= 2,
"GatherElements requires data and indices");
1008 MX data = densify(node_inputs[0]);
1009 MX indices_mx = node_inputs[1];
1011 MX idx_vec = vec(indices_mx);
1012 MX flat = vec(data);
1013 if (indices_mx.is_constant()) {
1018 flat.get_nz(output,
false, idx_vec);
1022 casadi_error(
"Unsupported operation '" + op_type +
"'");
casadi_int numel() const
Get the number of elements.
static MX sym(const std::string &name, casadi_int nrow=1, casadi_int ncol=1)
Create an nrow-by-ncol symbolic primitive.
bool verbose_
Verbose – for debugging.
static MX kron(const MX &x, const MX &b)
bool is_constant() const
Check if constant.
void get_nz(MX &m, bool ind1, const Slice &kk) const
Function create(const std::string &name)
Create a CasADi Function from the loaded ONNX graph.
onnx::ModelProto model_
ONNX model protocol buffer.
bool has_model_
Whether a model has been loaded.
template class CASADI_EXPORT Matrix< casadi_int >
static std::vector< casadi_int > constant_ints(const MX &m)
T norm_1(const std::vector< T > &x)
double if_else(double x, double y, double z)
double sign(double x)
Sign function, note that sign(nan) == nan.
double get_float_attribute(const onnx::NodeProto &node, const std::string &name, double default_value)
Read a float node attribute by name, or default_value if absent.
const onnx::GraphProto * get_graph_attribute(const onnx::NodeProto &node, const std::string &name)
Read a subgraph (GraphProto) node attribute by name, or nullptr if absent.
std::string str(const T &v)
String representation, any type.
static bool integer_tensor(const onnx::TensorProto &tensor, std::vector< int64_t > &values)
const OpMapping * get_op_mapping_by_name(const std::string &onnx_name)
Lookup operation mapping by ONNX name (for import)
casadi_int get_int_attribute(const onnx::NodeProto &node, const std::string &name, casadi_int default_value)
Read an integer node attribute by name, or default_value if absent.
T norm_2(const std::vector< T > &x)
std::string get_string_attribute(const onnx::NodeProto &node, const std::string &name)
Read a string node attribute by name, or "" if absent.