26 #include "onnx_model.hpp"
32 void Onnx::process_graph_initializers(
33 const onnx::GraphProto& graph,
34 std::map<std::string, MX>& value_map,
37 for (
int i = 0; i < graph.initializer_size(); ++i) {
38 const onnx::TensorProto& tensor = graph.initializer(i);
39 std::string tensor_name = tensor.name();
42 uout() <<
" Processing initializer: " << tensor_name << std::endl;
45 value_map[tensor_name] = MX(tensor_to_dm(tensor));
50 void Onnx::process_graph_inputs(
51 const onnx::GraphProto& graph,
52 std::map<std::string, MX>& value_map,
53 std::vector<MX>& func_inputs,
54 std::vector<std::string>& input_names,
57 for (
int i = 0; i < graph.input_size(); ++i) {
58 const onnx::ValueInfoProto& input = graph.input(i);
59 std::string input_name = input.name();
62 if (value_map.count(input_name)) {
64 uout() <<
" Skipping input '" << input_name
65 <<
"' (it's an initializer)" << std::endl;
71 const onnx::TensorShapeProto& shape =
72 input.type().tensor_type().shape();
73 casadi_int rows = get_dimension(shape, 0);
74 casadi_int cols = get_dimension(shape, 1);
77 uout() <<
" Creating input: " << input_name
78 <<
" [" << rows <<
", " << cols <<
"]" << std::endl;
86 Sparsity ov = input_pattern(graph, input_name);
87 MX mx_input = ov.is_null() ?
MX::sym(input_name, cols, rows) : MX::sym(input_name, ov);
88 value_map[input_name] = mx_input;
89 func_inputs.push_back(mx_input);
90 input_names.push_back(input_name);
95 void Onnx::process_graph_nodes(
96 const onnx::GraphProto& graph,
97 std::map<std::string, MX>& value_map,
101 std::map<std::string, MX> kron_operand_a_, kron_operand_b_;
103 for (
int i = 0; i < graph.node_size(); ++i) {
104 const onnx::NodeProto& node = graph.node(i);
105 std::string op_type = node.op_type();
108 uout() <<
" Processing node " << i <<
": " << op_type << std::endl;
121 if (node.name().rfind(
"kron", 0) == 0) {
122 const std::string& nm = node.name();
123 std::size_t us = nm.rfind(
'_');
124 std::string tag = nm.substr(0, us), role = nm.substr(us + 1);
126 kron_operand_a_[tag] = value_map.at(node.input(0));
130 kron_operand_b_[tag] = value_map.at(node.input(0));
133 if (role ==
"P")
continue;
136 casadi_assert(kron_operand_a_.count(tag) && kron_operand_b_.count(tag),
137 "ONNX import: incomplete kron node group '" + tag +
"'");
138 value_map[node.output(0)] =
MX::kron(kron_operand_a_.at(tag), kron_operand_b_.at(tag));
144 std::vector<MX> node_inputs;
145 for (
int j = 0; j < node.input_size(); ++j) {
146 std::string input_name = node.input(j);
149 if (input_name.empty()) {
150 node_inputs.push_back(MX());
154 casadi_assert(value_map.count(input_name),
155 "Unknown input tensor '" + input_name +
156 "' required by node " + std::to_string(i) +
157 " (op_type: " + op_type +
")");
158 node_inputs.push_back(value_map[input_name]);
165 if (op_type ==
"Split") {
166 casadi_assert(node_inputs.size() >= 1,
"Split requires 1 input");
168 casadi_assert(axis == 0 || axis == 1,
"Split: only axis 0 and 1 supported");
171 std::vector<casadi_int> split_sizes;
172 if (node_inputs.size() >= 2) {
173 DM split_dm =
DM(node_inputs[1]);
174 for (casadi_int k = 0; k < split_dm.nnz(); ++k) {
175 split_sizes.push_back(
static_cast<casadi_int
>(
static_cast<double>(split_dm.nz(k))));
178 for (
int a = 0; a < node.attribute_size(); ++a) {
179 if (node.attribute(a).name() ==
"split") {
180 for (
int k = 0; k < node.attribute(a).ints_size(); ++k) {
181 split_sizes.push_back(node.attribute(a).ints(k));
187 if (split_sizes.empty()) {
189 casadi_int total = (axis == 0) ? node_inputs[0].size2() : node_inputs[0].size1();
190 split_sizes.assign(node.output_size(), total / node.output_size());
194 std::vector<casadi_int> offset = {0};
195 for (casadi_int sz : split_sizes) offset.push_back(offset.back() + sz);
196 std::vector<MX> outputs = (axis == 0) ? horzsplit(node_inputs[0], offset)
197 : vertsplit(node_inputs[0], offset);
199 for (casadi_int j = 0; j < outputs.size(); ++j) value_map[node.output(j)] = outputs[j];
203 }
else if (op_type ==
"Scan") {
205 casadi_assert(body !=
nullptr,
"Scan node requires a 'body' subgraph");
206 casadi_int num_scan_inputs =
get_int_attribute(node,
"num_scan_inputs", node.input_size());
207 casadi_int M = node.input_size() - num_scan_inputs;
210 std::set<std::string> body_in_names, defined;
211 for (
int b = 0; b < body->input_size(); ++b) {
212 body_in_names.insert(body->input(b).name());
213 defined.insert(body->input(b).name());
215 for (
int nn = 0; nn < body->node_size(); ++nn)
216 for (
int oo = 0; oo < body->node(nn).output_size(); ++oo)
217 defined.insert(body->node(nn).output(oo));
218 bool has_capture =
false;
219 for (
int nn = 0; nn < body->node_size() && !has_capture; ++nn)
220 for (
int ii = 0; ii < body->node(nn).input_size(); ++ii) {
221 const std::string& in = body->node(nn).input(ii);
222 if (!in.empty() && !defined.count(in)) { has_capture =
true;
break; }
225 if (M == 0 && !has_capture) {
228 Function base = function_from_graph(*body, op_type +
"_body");
229 casadi_int n = node_inputs[0].size2() / base.size2_in(0);
230 std::vector<MX> outputs =
231 base.map(n)(std::vector<MX>(node_inputs.begin(), node_inputs.end()));
232 for (
int j = 0; j < node.output_size(); ++j) value_map[node.output(j)] = outputs[j];
238 const onnx::NodeProto* call =
nullptr;
239 for (
int nn = 0; nn < body->node_size(); ++nn)
240 if (!body->node(nn).domain().empty()) { call = &body->node(nn);
break; }
241 casadi_assert(call,
"reduce-map Scan body must contain a base function call");
242 casadi_int nin = call->input_size(), nout = call->output_size();
245 std::vector<bool> reduce_in(nin), reduce_out(nout,
false);
246 for (casadi_int j = 0; j < nin; ++j) reduce_in[j] = !body_in_names.count(call->input(j));
247 for (
int nn = 0; nn < body->node_size(); ++nn) {
248 if (body->node(nn).op_type() !=
"Add")
continue;
249 for (
int ii = 0; ii < body->node(nn).input_size(); ++ii)
250 for (casadi_int j = 0; j < nout; ++j)
251 if (body->node(nn).input(ii) == call->output(j)) reduce_out[j] =
true;
255 std::vector<std::pair<casadi_int, casadi_int>> in_shapes(nin);
256 for (casadi_int j = 0; j < nin; ++j) {
258 MX cap = value_map.at(call->input(j));
259 in_shapes[j] = {cap.size1(), cap.size2()};
261 for (
int b = 0; b < body->input_size(); ++b)
262 if (body->input(b).name() == call->input(j)) {
263 const onnx::TensorShapeProto& sh = body->input(b).type().tensor_type().shape();
264 in_shapes[j] = {get_dimension(sh, 0), get_dimension(sh, 1)};
270 const onnx::FunctionProto* fp = find_function(call->op_type(), call->domain());
271 casadi_assert(fp,
"reduce-map base function '" + call->op_type() +
"' not found");
272 Function base = function_from_function_proto(*fp, in_shapes, op_type +
"_base");
275 std::vector<MX> args(nin);
276 casadi_int n = 0, r = 0;
277 for (casadi_int j = 0; j < nin; ++j) {
279 args[j] = value_map.at(call->input(j));
281 args[j] = node_inputs[M + r];
282 if (n == 0) n = args[j].size2() / base.size2_in(j);
286 casadi_assert(n > 0,
"reduce-map import: could not infer map size");
287 std::vector<MX> outs = base.map(n, reduce_in, reduce_out)(args);
290 casadi_int si = 0, ci = 0;
291 for (casadi_int j = 0; j < nout; ++j) {
292 if (reduce_out[j]) value_map[node.output(si++)] = outs[j];
294 value_map[node.output(M + ci++)] = outs[j];
299 }
else if (op_type ==
"If") {
302 casadi_assert(then_b && else_b,
"If node requires then_branch and else_branch");
303 std::vector<MX> t = eval_captured_subgraph(*then_b, value_map);
304 std::vector<MX> e = eval_captured_subgraph(*else_b, value_map);
305 MX cond = node_inputs[0];
306 for (
int j = 0; j < node.output_size(); ++j) {
307 value_map[node.output(j)] =
if_else(cond, t[j], e[j]);
312 }
else if (op_type ==
"Loop") {
313 casadi_error(
"ONNX import: 'Loop' control flow operator is not supported.");
318 std::string node_domain = node.domain();
319 if (!node_domain.empty()) {
320 const onnx::FunctionProto* func_proto = find_function(op_type, node_domain);
322 if (func_proto !=
nullptr) {
324 uout() <<
" Function call to: " << node_domain <<
"." << op_type << std::endl;
329 onnx::GraphProto func_graph;
330 func_graph.set_name(func_proto->name());
331 for (
int n = 0; n < func_proto->node_size(); ++n) {
332 *func_graph.add_node() = func_proto->node(n);
336 std::map<std::string, MX> func_value_map;
337 for (
size_t n = 0; n < node_inputs.size() && n < func_proto->input_size(); ++n) {
338 func_value_map[func_proto->input(n)] = node_inputs[n];
341 process_graph_nodes(func_graph, func_value_map,
false);
344 for (
int n = 0; n < node.output_size() && n < func_proto->output_size(); ++n) {
345 std::string func_output_name = func_proto->output(n);
346 casadi_assert(func_value_map.count(func_output_name),
347 "Function output '" + func_output_name +
"' not found");
348 value_map[node.output(n)] = func_value_map[func_output_name];
351 uout() <<
" -> " << node.output(n) << std::endl;
362 output = process_node_operation(op_type, node, node_inputs);
366 casadi_assert(node.output_size() >= 1,
367 "Node must have at least one output");
369 std::string output_name = node.output(0);
370 value_map[output_name] = output;
373 uout() <<
" -> " << output_name << std::endl;
379 void Onnx::collect_graph_outputs(
380 const onnx::GraphProto& graph,
381 const std::map<std::string, MX>& value_map,
382 std::vector<MX>& func_outputs,
383 std::vector<std::string>& output_names,
384 bool verbose)
const {
386 for (
int i = 0; i < graph.output_size(); ++i) {
387 const onnx::ValueInfoProto& output = graph.output(i);
388 std::string output_name = output.name();
390 casadi_assert(value_map.count(output_name),
391 "Unknown output tensor: " + output_name +
392 ". This usually means the ONNX graph contains unsupported operations.");
396 func_outputs.push_back(value_map.at(output_name));
397 output_names.push_back(output_name);
400 uout() <<
" Graph output: " << output_name << std::endl;
405 const onnx::FunctionProto* Onnx::find_function(
const std::string& name,
406 const std::string& domain)
const {
407 for (
int f = 0; f <
model_.functions_size(); ++f)
408 if (
model_.functions(f).name() == name &&
model_.functions(f).domain() == domain)
409 return &
model_.functions(f);
414 casadi_assert(
has_model_,
"No ONNX model loaded. Call load() first.");
415 return function_from_graph(
model_.graph(), name);
418 Function Onnx::function_from_graph(
const onnx::GraphProto& graph,
419 const std::string& name) {
420 std::map<std::string, MX> value_map;
421 std::vector<MX> inputs, outputs;
422 std::vector<std::string> input_names, output_names;
425 uout() <<
"Building CasADi Function '" << name <<
"' from graph '" << graph.name()
426 <<
"' (" << graph.initializer_size() <<
" initializers, "
427 << graph.input_size() <<
" inputs, " << graph.output_size() <<
" outputs, "
428 << graph.node_size() <<
" nodes)" << std::endl;
431 process_graph_initializers(graph, value_map,
verbose_);
432 process_graph_inputs(graph, value_map, inputs, input_names,
verbose_);
433 process_graph_nodes(graph, value_map,
verbose_);
434 collect_graph_outputs(graph, value_map, outputs, output_names,
verbose_);
436 return Function(name, inputs, outputs, input_names, output_names);
439 Function Onnx::function_from_function_proto(
440 const onnx::FunctionProto& fp,
441 const std::vector<std::pair<casadi_int, casadi_int>>& in_shapes,
442 const std::string& name) {
445 g.set_name(fp.name());
446 for (
int j = 0; j < fp.input_size(); ++j) {
447 onnx::ValueInfoProto* vi = g.add_input();
448 vi->set_name(fp.input(j));
449 onnx::TypeProto::Tensor* tt = vi->mutable_type()->mutable_tensor_type();
450 tt->set_elem_type(real_type());
451 tt->mutable_shape()->add_dim()->set_dim_value(in_shapes[j].first);
452 tt->mutable_shape()->add_dim()->set_dim_value(in_shapes[j].second);
454 for (
int k = 0; k < fp.node_size(); ++k) *g.add_node() = fp.node(k);
455 for (
int j = 0; j < fp.output_size(); ++j) g.add_output()->set_name(fp.output(j));
456 return function_from_graph(g, name);
459 std::vector<MX> Onnx::eval_captured_subgraph(
const onnx::GraphProto& graph,
460 std::map<std::string, MX> scope) {
462 process_graph_initializers(graph, scope,
verbose_);
463 process_graph_nodes(graph, scope,
verbose_);
464 std::vector<MX> outputs;
465 for (
int i = 0; i < graph.output_size(); ++i)
466 outputs.push_back(scope.at(graph.output(i).name()));
473 casadi_assert(m.
is_constant(),
"Expected a constant integer tensor");
474 DM dm =
static_cast<DM>(m);
475 std::vector<casadi_int> v;
476 for (casadi_int k = 0; k < dm.
numel(); ++k)
477 v.push_back(
static_cast<casadi_int
>(dm(k).scalar()));
481 MX Onnx::process_node_operation(
482 const std::string& op_type,
483 const onnx::NodeProto& node,
484 const std::vector<MX>& node_inputs) {
491 const MX& x = node_inputs[0];
492 if (mapping->arity == 1) {
493 casadi_assert(node_inputs.size() >= 1, op_type +
" requires 1 input");
494 switch (mapping->casadi_op) {
495 case OP_SIN:
return sin(x);
496 case OP_COS:
return cos(x);
497 case OP_TAN:
return tan(x);
507 case OP_EXP:
return exp(x);
508 case OP_LOG:
return log(x);
515 case OP_ERF:
return erf(x);
516 case OP_INV:
return 1.0 / x;
525 }
else if (mapping->arity == 2) {
526 casadi_assert(node_inputs.size() >= 2, op_type +
" requires 2 inputs");
527 const MX& y = node_inputs[1];
528 switch (mapping->casadi_op) {
529 case OP_ADD:
return x + y;
530 case OP_SUB:
return x - y;
531 case OP_MUL:
return x * y;
532 case OP_DIV:
return x / y;
533 case OP_POW:
return pow(x, y);
541 if (op_type ==
"Less") {
542 casadi_assert(node_inputs.size() >= 2,
"Less requires 2 inputs");
543 output =
if_else(node_inputs[0] < node_inputs[1], MX(1.0), MX(0.0));
545 }
else if (op_type ==
"Equal") {
546 casadi_assert(node_inputs.size() >= 2,
"Equal requires 2 inputs");
547 output = !ne(node_inputs[0], node_inputs[1]);
549 }
else if (op_type ==
"LessOrEqual") {
550 casadi_assert(node_inputs.size() >= 2,
"LessOrEqual requires 2 inputs");
551 output =
if_else(node_inputs[0] <= node_inputs[1], MX(1.0), MX(0.0));
553 }
else if (op_type ==
"Min") {
554 casadi_assert(node_inputs.size() >= 2,
"Min requires 2 inputs");
555 output = fmin(node_inputs[0], node_inputs[1]);
557 }
else if (op_type ==
"Max") {
558 casadi_assert(node_inputs.size() >= 2,
"Max requires 2 inputs");
559 output = fmax(node_inputs[0], node_inputs[1]);
561 }
else if (op_type ==
"Mod") {
562 casadi_assert(node_inputs.size() >= 2,
"Mod requires 2 inputs");
563 output = fmod(node_inputs[0], node_inputs[1]);
565 }
else if (op_type ==
"ReduceSum") {
567 casadi_assert(node_inputs.size() >= 1,
"ReduceSum requires 1 input");
568 output = sum1(sum2(node_inputs[0]));
570 }
else if (op_type ==
"Not") {
571 casadi_assert(node_inputs.size() >= 1,
"Not requires 1 input");
572 output = logic_not(node_inputs[0]);
574 }
else if (op_type ==
"And") {
575 casadi_assert(node_inputs.size() >= 2,
"And requires 2 inputs");
576 output = logic_and(node_inputs[0], node_inputs[1]);
578 }
else if (op_type ==
"Or") {
579 casadi_assert(node_inputs.size() >= 2,
"Or requires 2 inputs");
580 output = logic_or(node_inputs[0], node_inputs[1]);
582 }
else if (op_type ==
"Where") {
583 casadi_assert(node_inputs.size() >= 3,
"Where requires 3 inputs");
584 output =
if_else(node_inputs[0], node_inputs[1], node_inputs[2]);
586 }
else if (op_type ==
"Identity") {
587 casadi_assert(node_inputs.size() >= 1,
"Identity requires 1 input");
588 output = node_inputs[0];
590 }
else if (op_type ==
"Cast") {
593 casadi_assert(node_inputs.size() >= 1,
"Cast requires 1 input");
594 output = node_inputs[0];
596 }
else if (op_type ==
"MatMul") {
598 casadi_assert(node_inputs.size() >= 2,
"MatMul requires 2 inputs");
599 output = mtimes(node_inputs[1], node_inputs[0]);
601 }
else if (op_type ==
"Gemm") {
604 casadi_assert(node_inputs.size() >= 2,
"Gemm requires at least 2 inputs");
605 MX A =
get_int_attribute(node,
"transB", 0) ? node_inputs[1].T() : node_inputs[1];
606 MX B =
get_int_attribute(node,
"transA", 0) ? node_inputs[0].T() : node_inputs[0];
608 if (node_inputs.size() >= 3) {
612 }
else if (op_type ==
"Sum") {
614 casadi_assert(node_inputs.size() >= 1,
"Sum requires at least 1 input");
615 output = node_inputs[0];
616 for (casadi_int idx = 1; idx < node_inputs.size(); ++idx) output = output + node_inputs[idx];
618 }
else if (op_type ==
"Pad") {
622 casadi_assert(node_inputs.size() >= 2,
"Pad requires data and pads");
623 MX block = node_inputs[0];
624 std::vector<casadi_int> pads =
constant_ints(node_inputs[1]);
625 casadi_int col_off = pads[0], row_off = pads[1];
626 casadi_int br = block.size1(), bc = block.size2();
627 casadi_int
C = col_off + bc + pads[2];
632 if (col_off > 0 || pads[2] > 0) {
633 padded = horzcat(MX(Sparsity(br, col_off)), padded, MX(Sparsity(br, pads[2])));
635 if (row_off > 0 || pads[3] > 0) {
636 padded = vertcat(MX(Sparsity(row_off, C)), padded, MX(Sparsity(pads[3], C)));
640 }
else if (op_type ==
"Einsum") {
643 casadi_assert(node_inputs.size() >= 2,
"Einsum requires 2 inputs");
646 size_t comma = eq.find(
','), arrow = eq.find(
"->");
647 casadi_assert(comma != std::string::npos && arrow != std::string::npos,
648 "ONNX import: only binary Einsum 'a,b->c' is supported");
649 std::string sa = eq.substr(0, comma);
650 std::string sb = eq.substr(comma + 1, arrow - comma - 1);
651 std::string sc = eq.substr(arrow + 2);
657 std::reverse(sa.begin(), sa.end());
658 std::reverse(sb.begin(), sb.end());
659 std::reverse(sc.begin(), sc.end());
665 std::map<char, casadi_int> lsize;
666 for (
int t = 0; t < 2; ++t) {
667 const std::string& s = (t == 0) ? sa : sb;
668 const MX& m = node_inputs[t];
670 lsize[s[0]] = m.numel();
671 }
else if (s.size() >= 2) {
672 lsize[s[0]] = m.size1();
673 lsize[s[1]] = m.size2();
677 std::map<char, casadi_int> lab;
678 casadi_int next_label = -1;
679 for (
char ch : sa + sb + sc)
if (!lab.count(ch)) lab[ch] = next_label--;
682 std::vector<casadi_int> da, db, dc, La, Lb, Lc;
683 for (
char ch : sa) { da.push_back(lsize[ch]); La.push_back(lab[ch]); }
684 for (
char ch : sb) { db.push_back(lsize[ch]); Lb.push_back(lab[ch]); }
685 for (
char ch : sc) { dc.push_back(lsize[ch]); Lc.push_back(lab[ch]); }
686 output = einstein(vec(node_inputs[0]), vec(node_inputs[1]), da, db, dc, La, Lb, Lc);
688 }
else if (op_type ==
"Det") {
689 casadi_assert(node_inputs.size() >= 1,
"Det requires 1 input");
690 output = det(node_inputs[0]);
692 }
else if (op_type ==
"ReduceLogSumExp") {
694 casadi_assert(node_inputs.size() >= 1,
"ReduceLogSumExp requires 1 input");
695 output = log(sum1(sum2(exp(node_inputs[0]))));
697 }
else if (op_type ==
"Constant") {
699 const onnx::AttributeProto* value_attr =
nullptr;
700 const onnx::AttributeProto* sparse_attr =
nullptr;
701 for (
int a = 0; a < node.attribute_size(); ++a) {
702 const std::string& an = node.attribute(a).name();
703 if (an ==
"value") value_attr = &node.attribute(a);
704 else if (an ==
"sparse_value") sparse_attr = &node.attribute(a);
706 if (sparse_attr !=
nullptr) {
707 output = MX(sparse_tensor_to_dm(sparse_attr->sparse_tensor()));
709 casadi_assert(value_attr !=
nullptr,
710 "Constant node must have a 'value' or 'sparse_value' attribute");
711 output = MX(tensor_to_dm(value_attr->t()));
715 }
else if (op_type ==
"Reshape") {
716 casadi_assert(node_inputs.size() >= 2,
717 "Reshape operation requires 2 inputs (data and shape)");
719 casadi_assert(node_inputs[1].is_constant(),
720 "Reshape shape must be a constant");
721 DM shape_dm =
static_cast<DM>(node_inputs[1]);
727 if (shape_dm.numel() > 2) {
729 output = densify(node_inputs[0]);
732 casadi_int s0 =
static_cast<casadi_int
>(shape_dm(0).scalar());
733 casadi_int s1 = (shape_dm.numel() > 1) ?
static_cast<casadi_int
>(shape_dm(1).scalar()) : 1;
734 output = reshape(densify(node_inputs[0]), s1, s0);
737 }
else if (op_type ==
"Concat") {
741 output = horzcat(node_inputs);
742 }
else if (axis == 1) {
743 output = vertcat(node_inputs);
745 casadi_error(
"Concat with axis=" + std::to_string(axis) +
746 " not supported. Only axis=0 (vertcat) and axis=1 (horzcat) are supported.");
749 }
else if (op_type ==
"Slice") {
751 casadi_assert(node_inputs.size() >= 3,
752 "Slice requires at least 3 inputs (data, starts, ends)");
754 MX data = node_inputs[0];
755 casadi_assert(node_inputs[1].is_constant() && node_inputs[2].is_constant(),
756 "Slice starts and ends must be constants");
758 DM starts_dm =
static_cast<DM>(node_inputs[1]);
759 DM ends_dm =
static_cast<DM>(node_inputs[2]);
762 std::vector<casadi_int> axes, steps;
763 if (node_inputs.size() >= 4 && !node_inputs[3].is_empty()) {
766 for (casadi_int k = 0; k < starts_dm.numel(); ++k) axes.push_back(k);
768 if (node_inputs.size() >= 5 && !node_inputs[4].is_empty()) {
771 steps.assign(starts_dm.numel(), 1);
774 casadi_assert(axes.size() <= 2,
"Slice: only up to 2D slicing supported");
778 Slice row_slice, col_slice;
779 casadi_int nrow = data.size1(), ncol = data.size2();
780 for (casadi_int a = 0; a < static_cast<casadi_int>(axes.size()); ++a) {
781 casadi_int st =
static_cast<casadi_int
>(starts_dm(a).scalar());
782 casadi_int en =
static_cast<casadi_int
>(ends_dm(a).scalar());
783 casadi_int sp = steps[a];
785 col_slice = Slice(st, en > ncol ? ncol : en, sp);
787 row_slice = Slice(st, en > nrow ? nrow : en, sp);
790 output = data(row_slice, col_slice);
792 }
else if (op_type ==
"Gather") {
794 casadi_assert(node_inputs.size() >= 2,
"Gather requires data and indices");
797 MX data = densify(node_inputs[0]);
798 MX indices_mx = node_inputs[1];
801 casadi_assert(indices_mx.is_constant(),
"Gather indices must be constant");
802 DM indices_dm =
static_cast<DM>(indices_mx);
804 if (data.size2() == 1) {
807 if (indices_dm.numel() == 1) {
808 output = data(
static_cast<casadi_int
>(indices_dm(0).scalar()), Slice());
812 }
else if (indices_dm.numel() == 1) {
813 casadi_int idx =
static_cast<casadi_int
>(indices_dm(0).scalar());
815 output = data(idx, Slice());
816 }
else if (axis == 1) {
817 output = data(Slice(), idx);
819 casadi_error(
"Gather: only axis 0 and 1 supported for 2D tensors");
825 std::vector<MX> rows;
826 for (casadi_int idx : indices) rows.push_back(data(idx, Slice()));
827 output = vertcat(rows);
828 }
else if (axis == 1) {
829 std::vector<MX> cols;
830 for (casadi_int idx : indices) cols.push_back(data(Slice(), idx));
831 output = horzcat(cols);
833 casadi_error(
"Gather: only axis 0 and 1 supported for 2D tensors");
837 }
else if (op_type ==
"ScatterElements") {
840 casadi_assert(node_inputs.size() >= 3,
"ScatterElements requires data, indices, updates");
841 MX se_data = densify(node_inputs[0]);
842 MX se_idx = node_inputs[1];
843 MX se_upd = vec(node_inputs[2]);
849 MX w =
MX::sym(
"scel_w", se_data.numel(), 1), g;
850 if (se_idx.is_constant()) {
853 w.get_nz(g,
false, vec(se_idx));
855 output = se_data + reshape(jtimes(g, w, se_upd,
true), se_data.size1(), se_data.size2());
858 if (se_idx.is_constant()) {
861 output.set_nz(se_upd,
false, vec(se_idx));
865 }
else if (op_type ==
"ScatterND") {
869 casadi_assert(node_inputs.size() >= 3,
"ScatterND requires data, indices, updates");
870 MX data = densify(node_inputs[0]);
871 std::vector<casadi_int> coords =
constant_ints(node_inputs[1]);
872 casadi_int R = data.size1();
873 std::vector<casadi_int> pos(coords.size() / 2);
874 for (casadi_int i = 0; i < static_cast<casadi_int>(pos.size()); ++i) {
875 pos[i] = coords[2 * i] * R + coords[2 * i + 1];
882 MX w =
MX::sym(
"addnz_w", data.numel(), 1), g;
884 MX scatter = jtimes(g, w, vec(node_inputs[2]),
true);
885 output = data + reshape(scatter, data.size1(), data.size2());
891 }
else if (op_type ==
"Tile") {
893 casadi_assert(node_inputs.size() >= 2,
"Tile requires data and repeats inputs");
894 MX data = node_inputs[0];
895 MX repeats_mx = node_inputs[1];
896 casadi_assert(repeats_mx.is_constant(),
"Tile repeats must be constant");
897 DM repeats_dm =
static_cast<DM>(repeats_mx);
899 casadi_int rows_repeat = 1, cols_repeat = 1;
900 if (repeats_dm.numel() >= 1) rows_repeat =
static_cast<casadi_int
>(repeats_dm(0).scalar());
901 if (repeats_dm.numel() >= 2) cols_repeat =
static_cast<casadi_int
>(repeats_dm(1).scalar());
904 output = repmat(data, cols_repeat, rows_repeat);
906 }
else if (op_type ==
"GatherElements") {
911 casadi_assert(node_inputs.size() >= 2,
"GatherElements requires data and indices");
912 MX data = densify(node_inputs[0]);
913 MX indices_mx = node_inputs[1];
915 MX idx_vec = vec(indices_mx);
917 if (indices_mx.is_constant()) {
922 flat.get_nz(output,
false, idx_vec);
926 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.
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.