26 #include "onnx_model.hpp"
27 #include <casadi/core/casadi_meta.hpp>
36 return n.empty() ?
"input_" + std::to_string(i) : n;
41 return n.empty() ?
"output_" + std::to_string(i) : n;
44 void Onnx::set_real_tensor_type(onnx::ValueInfoProto* value,
const Sparsity& sp) {
47 onnx::TypeProto::Tensor* tensor_type = value->mutable_type()->mutable_tensor_type();
48 tensor_type->set_elem_type(real_type());
49 onnx::TensorShapeProto* shape = tensor_type->mutable_shape();
50 shape->add_dim()->set_dim_value(sp.size2());
51 shape->add_dim()->set_dim_value(sp.size1());
56 void Onnx::add_graph_inputs(onnx::GraphProto* graph,
const Function& f,
57 const std::string& name_prefix) {
58 for (casadi_int i = 0; i < f.n_in(); ++i) {
59 onnx::ValueInfoProto* input = graph->add_input();
61 : name_prefix +
"_i" + std::to_string(i));
62 set_real_tensor_type(input, f.sparsity_in(i));
67 void Onnx::add_graph_outputs(onnx::GraphProto* graph,
const Function& f,
68 const std::string& name_prefix) {
69 for (casadi_int i = 0; i < f.n_out(); ++i) {
70 onnx::ValueInfoProto* output = graph->add_output();
72 : name_prefix +
"_o" + std::to_string(i));
73 set_real_tensor_type(output, f.sparsity_out(i));
81 model_.set_producer_name(
"CasADi");
84 onnx::OperatorSetIdProto* opset =
model_.add_opset_import();
85 opset->set_domain(
"");
86 opset->set_version(16);
88 onnx::GraphProto* graph =
model_.mutable_graph();
89 graph->set_name(f.
name());
90 add_graph_inputs(graph, f);
93 std::map<casadi_int, std::string> work_to_onnx;
98 std::map<casadi_int, casadi_int> output_segment_count;
99 for (casadi_int k = 0; k < n_instr; ++k) {
103 casadi_int output_idx = info[
"ind"];
104 output_segment_count[output_idx]++;
109 std::map<casadi_int, std::map<casadi_int, std::string>> output_segment_values;
111 std::map<casadi_int, std::map<casadi_int, Sparsity>> output_segment_sparsity;
113 for (casadi_int k = 0; k < n_instr; ++k) {
119 std::string node_output =
"n" + std::to_string(k);
125 casadi_int output_idx = info[
"ind"];
126 casadi_int offset = info[
"offset"];
129 std::string input_onnx_name = work_to_onnx[i[0]];
131 if (output_segment_count[output_idx] > 1) {
135 output_segment_values[output_idx][offset] = input_onnx_name;
136 output_segment_sparsity[output_idx][offset] = mx.
dep(0).
sparsity();
145 auto add_node = [graph]() -> onnx::NodeProto* {
return graph->add_node(); };
148 std::string tmp =
"out_pre_" + std::to_string(k);
149 emit_output_node(graph, input_onnx_name, dep.
size1(), dep.
size2(), out_sp,
150 tmp,
"out_rs_" + std::to_string(k));
151 emit_sparsity_restore(add_node, tmp,
153 "outr_" + std::to_string(k), output_name);
155 emit_output_node(graph, input_onnx_name, dep.
size1(), dep.
size2(), out_sp,
156 output_name,
"out_rs_" + std::to_string(k));
162 if (process_operation(graph, f, op, k, i, o, work_to_onnx, node_output))
continue;
167 if (is_map_function(called_func)) {
168 export_map(graph, called_func, i, o, work_to_onnx, node_output +
"_out");
169 }
else if (is_reduce_map_function(called_func)) {
170 export_reduce_map(graph, called_func, i, o, work_to_onnx, node_output +
"_out");
171 }
else if (is_if_else_function(called_func)) {
172 export_if(graph, called_func, i, o, work_to_onnx, node_output +
"_out");
174 assert_not_control_flow(called_func);
175 export_call(graph, called_func, i, o, work_to_onnx, node_output +
"_out");
181 casadi_error(
"ONNX export: unsupported operation code " +
182 std::to_string(op) +
" at instruction " + std::to_string(k) +
183 ". The CasADi Function contains operations that cannot be exported to ONNX.");
196 auto seg_add_node = [graph]() -> onnx::NodeProto* {
return graph->add_node(); };
197 for (
const auto& kv : output_segment_values) {
198 casadi_int output_idx = kv.first;
199 const auto& segments = kv.second;
200 const auto& seg_sp = output_segment_sparsity[output_idx];
203 std::vector<std::string> names;
204 std::vector<casadi_int> brs, bcs;
205 bool all_full_height =
true, all_full_width =
true;
206 for (
const auto& seg : segments) {
207 const Sparsity& sp = seg_sp.at(seg.first);
208 names.push_back(seg.second);
209 brs.push_back(sp.
size1()); bcs.push_back(sp.
size2());
210 if (sp.
size1() != R) all_full_height =
false;
211 if (sp.
size2() !=
C) all_full_width =
false;
215 std::string uniq =
"oseg" + std::to_string(output_idx);
218 auto finish = [&](
const std::function<void(
const std::string&)>& emit) ->
void {
222 std::string tmp = oname +
"_pre";
229 if (all_full_height) {
231 finish([&](
const std::string& dst) {
232 onnx::NodeProto* cc = seg_add_node();
233 cc->set_op_type(
"Concat");
234 for (
const auto& n : names) cc->add_input(n);
238 }
else if (all_full_width) {
240 finish([&](
const std::string& dst) {
241 onnx::NodeProto* cc = seg_add_node();
242 cc->set_op_type(
"Concat");
243 for (
const auto& n : names) cc->add_input(n);
249 std::vector<casadi_int> row_off, col_off;
250 casadi_int ro = 0, co = 0;
251 for (casadi_int b = 0; b < static_cast<casadi_int>(names.size()); ++b) {
252 row_off.push_back(ro); col_off.push_back(co);
253 ro += brs[b]; co += bcs[b];
255 finish([&](
const std::string& dst) {
256 emit_blockdiag(seg_add_node, names, row_off, col_off, brs, bcs, R,
C, dst, uniq);
268 std::set<std::string> output_names;
271 std::map<std::string, int> consumers;
272 for (
const auto& nd : graph->node())
273 for (
const auto& in : nd.input()) consumers[in]++;
275 std::map<std::string, int> producers;
276 for (
const auto& nd : graph->node())
277 for (
const auto& o : nd.output()) producers[o]++;
279 auto* nodes = graph->mutable_node();
281 std::map<std::string, std::string> rename;
283 for (
int n = 0; n < nodes->size(); ++n) {
284 const onnx::NodeProto& nd = nodes->Get(n);
285 if (nd.op_type() !=
"Identity" || nd.input_size() != 1 || nd.output_size() != 1)
continue;
286 const std::string& src = nd.input(0);
287 const std::string& dst = nd.output(0);
288 if (!output_names.count(dst))
continue;
289 if (output_names.count(src))
continue;
290 if (producers[src] != 1)
continue;
291 if (consumers[src] != 1)
continue;
292 if (rename.count(src) || drop.count(n))
continue;
298 for (
int n = 0; n < nodes->size(); ++n) {
299 if (drop.count(n))
continue;
300 onnx::NodeProto* nd = nodes->Mutable(n);
301 for (
int o = 0; o < nd->output_size(); ++o) {
302 auto it = rename.find(nd->output(o));
303 if (it != rename.end()) nd->set_output(o, it->second);
306 google::protobuf::RepeatedPtrField<onnx::NodeProto> kept;
307 for (
int n = 0; n < nodes->size(); ++n)
308 if (!drop.count(n)) kept.Add()->CopyFrom(nodes->Get(n));
314 add_graph_outputs(graph, f,
"");
321 std::set<std::string> produced;
322 for (
const auto& nd : graph->node())
323 for (
const auto& o : nd.output()) produced.insert(o);
324 auto add_node = [graph]() -> onnx::NodeProto* {
return graph->add_node(); };
325 for (casadi_int i = 0; i < f.
n_out(); ++i) {
327 if (!produced.count(oname)) {
330 add_real_constant(add_node, oname,
331 std::vector<double>(sp.
size1() * sp.
size2(), 0.0),
332 {sp.size2(), sp.size1()});
334 add_sparse_constant(add_node, oname,
DM(sp, 0.0));
346 for (casadi_int i = 0; i < f.
n_in(); ++i) {
348 fill_sparse_tensor(graph->add_sparse_initializer(),
onnx_input_name(f, i),
355 onnx::OperatorSetIdProto* casadi_opset =
model_.add_opset_import();
356 casadi_opset->set_domain(
"casadi");
357 casadi_opset->set_version(1);
361 <<
" function(s) to casadi domain" << std::endl;
368 uout() <<
"Converted CasADi Function to ONNX model: " << f.
name() << std::endl;
369 uout() <<
" Instructions processed: " << n_instr << std::endl;
370 uout() <<
" ONNX nodes created: " << graph->node_size() << std::endl;
376 bool Onnx::is_if_else_function(
const Function& f)
const {
381 bool Onnx::is_mapaccum_function(
const Function& f)
const {
384 std::string fname = f.name();
385 return fname.find(
"mapaccum") != std::string::npos ||
386 fname.find(
"accum") != std::string::npos;
389 bool Onnx::is_map_function(
const Function& f)
const {
392 std::string fname = f.name();
393 return fname.find(
"mapaccum") == std::string::npos &&
394 fname.find(
"accum") == std::string::npos;
397 bool Onnx::is_reduce_map_function(
const Function& f)
const {
399 for (
const std::string& nm : f.get_function()) {
400 if (f.get_function(nm).class_name() ==
"MapSum")
return true;
405 void Onnx::assert_not_control_flow(
const Function& called_func)
const {
406 casadi_assert(!is_mapaccum_function(called_func),
407 "ONNX export: mapaccum (Loop) is not supported.");
414 std::set<std::string> defined;
415 for (
int i = 0; i < g->input_size(); ++i) defined.insert(g->input(i).name());
416 for (
int i = 0; i < g->initializer_size(); ++i) defined.insert(g->initializer(i).name());
417 for (
int n = 0; n < g->node_size(); ++n) {
418 for (
int j = 0; j < g->node(n).output_size(); ++j) defined.insert(g->node(n).output(j));
420 for (
int n = 0; n < g->node_size(); ++n) {
421 onnx::NodeProto* nd = g->mutable_node(n);
422 for (
int i = 0; i < nd->input_size(); ++i) {
423 if (defined.count(nd->input(i))) nd->set_input(i, prefix + nd->input(i));
425 for (
int i = 0; i < nd->output_size(); ++i) nd->set_output(i, prefix + nd->output(i));
427 for (
int i = 0; i < g->input_size(); ++i) {
428 g->mutable_input(i)->set_name(prefix + g->input(i).name());
430 for (
int i = 0; i < g->output_size(); ++i) {
431 g->mutable_output(i)->set_name(prefix + g->output(i).name());
435 template<
typename Container>
436 void Onnx::emit_reshape(Container* container,
const std::string& data,
437 const std::vector<casadi_int>& shape,
438 const std::string& output,
const std::string& shape_name) {
439 onnx::NodeProto* sc = container->add_node();
440 sc->set_op_type(
"Constant");
441 sc->add_output(shape_name);
442 onnx::AttributeProto* attr = sc->add_attribute();
443 attr->set_name(
"value");
444 attr->set_type(onnx::AttributeProto::TENSOR);
445 onnx::TensorProto* t = attr->mutable_t();
446 t->set_data_type(onnx::TensorProto::INT64);
447 t->add_dims(
static_cast<casadi_int
>(shape.size()));
448 for (casadi_int v : shape) t->add_int64_data(v);
449 onnx::NodeProto* r = container->add_node();
450 r->set_op_type(
"Reshape");
452 r->add_input(shape_name);
453 r->add_output(output);
456 template<
typename Container>
457 void Onnx::colmajor_reshape_into(Container* container,
const std::string& data,
458 const std::vector<casadi_int>& dims,
459 const std::string& output,
const std::string& uniq) {
462 std::vector<casadi_int> rev(dims.rbegin(), dims.rend());
463 emit_reshape(container, data, rev, output, uniq +
"_s");
466 template<
typename Container>
467 void Onnx::emit_output_node(Container* container,
const std::string& data,
468 casadi_int src_rows, casadi_int src_cols,
const Sparsity& out_sp,
469 const std::string& output,
const std::string& uniq) {
471 if (src_rows != out_sp.size1() || src_cols != out_sp.size2()) {
472 colmajor_reshape_into(container, data, {out_sp.size1(), out_sp.size2()}, output, uniq);
474 onnx::NodeProto*
id = container->add_node();
475 id->set_op_type(
"Identity");
477 id->add_output(output);
481 onnx::GraphProto Onnx::build_scan_body(
const Function& base) {
484 onnx::GraphProto body;
485 body.set_name(base.name() +
"_scan_body");
486 std::map<casadi_int, std::string> work_to_onnx;
488 for (casadi_int j = 0; j < base.n_in(); ++j) {
489 onnx::ValueInfoProto* vi = body.add_input();
490 vi->set_name(
"body_in_" + std::to_string(j));
491 set_real_tensor_type(vi, base.sparsity_in(j));
494 casadi_int n_instr = base.n_instructions();
495 for (casadi_int k = 0; k < n_instr; ++k) {
496 casadi_int op = base.instruction_id(k);
497 std::vector<casadi_int> o = base.instruction_output(k);
498 std::vector<casadi_int> i_vec = base.instruction_input(k);
499 std::string node_output =
"n" + std::to_string(k);
502 work_to_onnx[o[0]] =
"body_in_" + std::to_string(i_vec[0]);
507 std::string out_name =
"body_out_" + std::to_string(o[0]);
508 MX dep = base.instruction_MX(k).dep(0);
509 Sparsity out_sp = base.sparsity_out(o[0]);
510 emit_output_node(&body, work_to_onnx[i_vec[0]], dep.size1(), dep.size2(), out_sp,
511 out_name,
"body_out_rs_" + std::to_string(k));
512 onnx::ValueInfoProto* vi = body.add_output();
513 vi->set_name(out_name);
514 set_real_tensor_type(vi, out_sp);
519 Function called = base.instruction_MX(k).which_function();
520 if (is_map_function(called)) {
521 export_map(&body, called, i_vec, o, work_to_onnx, node_output +
"_out");
522 }
else if (is_reduce_map_function(called)) {
523 export_reduce_map(&body, called, i_vec, o, work_to_onnx, node_output +
"_out");
524 }
else if (is_if_else_function(called)) {
525 export_if(&body, called, i_vec, o, work_to_onnx, node_output +
"_out");
527 assert_not_control_flow(called);
528 export_call(&body, called, i_vec, o, work_to_onnx, node_output +
"_out");
533 if (process_operation(&body, base, op, k, i_vec, o, work_to_onnx, node_output))
continue;
535 casadi_error(
"ONNX export: unsupported operation code " + std::to_string(op) +
536 " in Map body of '" + base.name() +
"'");
541 template<
typename Container>
542 void Onnx::export_map(Container* container,
const Function& map_fn,
543 const std::vector<casadi_int>& i_vec,
544 const std::vector<casadi_int>& o,
545 std::map<casadi_int, std::string>& work_to_onnx,
546 const std::string& out_prefix) {
547 Function base = map_fn.get_function(map_fn.get_function().at(0));
548 casadi_int n = map_fn.size2_out(0) / base.size2_out(0);
552 for (casadi_int j = 0; j < base.n_in(); ++j) {
553 casadi_assert(map_fn.size2_in(j) == n * base.size2_in(j),
554 "ONNX export: Map with non-repeated (reduce_in) inputs is not yet supported.");
556 for (casadi_int j = 0; j < base.n_out(); ++j) {
557 casadi_assert(map_fn.size2_out(j) == n * base.size2_out(j),
558 "ONNX export: Map with reduced (reduce_out) outputs is not yet supported.");
563 std::vector<std::string> scan_inputs;
564 for (casadi_int j = 0; j < base.n_in(); ++j) {
565 std::string lifted = out_prefix +
"_in" + std::to_string(j);
566 emit_reshape(container, work_to_onnx[i_vec[j]],
567 {n, base.size2_in(j), base.size1_in(j)}, lifted,
568 out_prefix +
"_insh" + std::to_string(j));
569 scan_inputs.push_back(lifted);
572 onnx::NodeProto* scan = container->add_node();
573 scan->set_op_type(
"Scan");
574 for (
const std::string& s : scan_inputs) scan->add_input(s);
575 std::vector<std::string> scan_outputs;
576 for (casadi_int j = 0; j < base.n_out(); ++j) {
577 std::string so = out_prefix +
"_s" + std::to_string(j);
578 scan->add_output(so);
579 scan_outputs.push_back(so);
583 add_ints_attribute(scan,
"scan_output_axes", std::vector<casadi_int>(base.n_out(), 0));
585 onnx::GraphProto body = build_scan_body(base);
588 onnx::AttributeProto* body_attr = scan->add_attribute();
589 body_attr->set_name(
"body");
590 body_attr->set_type(onnx::AttributeProto::GRAPH);
591 *body_attr->mutable_g() = body;
594 for (casadi_int j = 0; j < base.n_out(); ++j) {
595 std::string out = out_prefix + std::to_string(j);
596 emit_reshape(container, scan_outputs[j],
597 {n * base.size2_out(j), base.size1_out(j)}, out,
598 out_prefix +
"_outsh" + std::to_string(j));
599 work_to_onnx[o[j]] = out;
603 onnx::GraphProto Onnx::build_reduce_scan_body(
const Function& base,
604 const std::vector<bool>& reduce_in,
const std::vector<bool>& reduce_out,
605 const std::vector<std::string>& capture_names) {
608 onnx::GraphProto body;
609 body.set_name(base.name() +
"_redscan_body");
610 std::map<casadi_int, std::string> w;
612 for (casadi_int j = 0; j < base.n_out(); ++j) {
613 if (!reduce_out[j])
continue;
614 onnx::ValueInfoProto* vi = body.add_input();
615 vi->set_name(
"acc_in_" + std::to_string(j));
616 set_real_tensor_type(vi, base.sparsity_out(j));
619 std::vector<std::string> args(base.n_in());
620 for (casadi_int j = 0; j < base.n_in(); ++j) {
622 args[j] = capture_names[j];
624 args[j] =
"scan_in_" + std::to_string(j);
625 onnx::ValueInfoProto* vi = body.add_input();
626 vi->set_name(args[j]);
627 set_real_tensor_type(vi, base.sparsity_in(j));
631 std::vector<casadi_int> ii(base.n_in()), oo(base.n_out());
632 for (casadi_int j = 0; j < base.n_in(); ++j) { ii[j] = j; w[j] = args[j]; }
633 for (casadi_int j = 0; j < base.n_out(); ++j) oo[j] = base.n_in() + j;
634 export_call(&body, base, ii, oo, w,
"bcall_");
636 for (casadi_int j = 0; j < base.n_out(); ++j) {
637 if (!reduce_out[j])
continue;
638 std::string out =
"acc_out_" + std::to_string(j);
640 onnx::ValueInfoProto* vi = body.add_output();
642 set_real_tensor_type(vi, base.sparsity_out(j));
644 for (casadi_int j = 0; j < base.n_out(); ++j) {
645 if (reduce_out[j])
continue;
646 std::string out =
"yscan_" + std::to_string(j);
648 onnx::ValueInfoProto* vi = body.add_output();
650 set_real_tensor_type(vi, base.sparsity_out(j));
655 template<
typename Container>
656 void Onnx::export_reduce_map(Container* container,
const Function& wrapper,
657 const std::vector<casadi_int>& i_vec,
658 const std::vector<casadi_int>& o,
659 std::map<casadi_int, std::string>& work_to_onnx,
660 const std::string& out_prefix) {
663 for (
const std::string& nm : wrapper.get_function()) {
664 if (wrapper.get_function(nm).class_name() ==
"MapSum") {
665 mapsum = wrapper.get_function(nm);
669 Function base = mapsum.get_function(mapsum.get_function().at(0));
672 std::vector<bool> reduce_in(base.n_in()), reduce_out(base.n_out());
673 for (casadi_int j = 0; j < base.n_in(); ++j)
674 reduce_in[j] = wrapper.size2_in(j) == base.size2_in(j);
675 for (casadi_int j = 0; j < base.n_out(); ++j)
676 reduce_out[j] = wrapper.size2_out(j) == base.size2_out(j);
678 for (casadi_int j = 0; j < base.n_in(); ++j)
679 if (!reduce_in[j]) { n = wrapper.size2_in(j) / base.size2_in(j);
break; }
680 if (n == 0)
for (casadi_int j = 0; j < base.n_out(); ++j)
681 if (!reduce_out[j]) { n = wrapper.size2_out(j) / base.size2_out(j);
break; }
682 casadi_assert(n > 0,
"ONNX export: reduce-map with no repeated inputs or outputs.");
684 AddNodeFn add_node = [container]() -> onnx::NodeProto* {
return container->add_node(); };
687 std::vector<std::string> capture_names(base.n_in());
688 for (casadi_int j = 0; j < base.n_in(); ++j) {
689 if (!reduce_in[j])
continue;
690 capture_names[j] = out_prefix +
"_cap" + std::to_string(j);
691 create_unary_node(add_node,
"Identity", work_to_onnx[i_vec[j]], capture_names[j]);
695 std::vector<std::string> scan_inputs;
696 for (casadi_int j = 0; j < base.n_in(); ++j) {
697 if (reduce_in[j])
continue;
698 std::string lifted = out_prefix +
"_in" + std::to_string(j);
699 emit_reshape(container, work_to_onnx[i_vec[j]],
700 {n, base.size2_in(j), base.size1_in(j)}, lifted,
701 out_prefix +
"_insh" + std::to_string(j));
702 scan_inputs.push_back(lifted);
706 std::vector<std::string> acc_inits;
707 for (casadi_int j = 0; j < base.n_out(); ++j) {
708 if (!reduce_out[j])
continue;
709 std::string nm = out_prefix +
"_acc0_" + std::to_string(j);
710 std::vector<double> zeros(base.size1_out(j) * base.size2_out(j), 0.0);
711 add_real_constant(add_node, nm, zeros, {base.size2_out(j), base.size1_out(j)});
712 acc_inits.push_back(nm);
715 onnx::NodeProto* scan = add_node();
716 scan->set_op_type(
"Scan");
717 for (
const std::string& s : acc_inits) scan->add_input(s);
718 for (
const std::string& s : scan_inputs) scan->add_input(s);
721 std::vector<std::string> state_out, scan_out;
722 for (casadi_int j = 0; j < base.n_out(); ++j)
724 state_out.push_back(out_prefix +
"_acc" + std::to_string(j));
725 scan->add_output(state_out.back());
727 for (casadi_int j = 0; j < base.n_out(); ++j)
728 if (!reduce_out[j]) {
729 scan_out.push_back(out_prefix +
"_s" + std::to_string(j));
730 scan->add_output(scan_out.back());
733 add_int_attribute(scan,
"num_scan_inputs",
static_cast<casadi_int
>(scan_inputs.size()));
734 add_ints_attribute(scan,
"scan_input_axes", std::vector<casadi_int>(scan_inputs.size(), 0));
735 add_ints_attribute(scan,
"scan_output_axes", std::vector<casadi_int>(scan_out.size(), 0));
737 onnx::GraphProto body = build_reduce_scan_body(base, reduce_in, reduce_out, capture_names);
739 onnx::AttributeProto* battr = scan->add_attribute();
740 battr->set_name(
"body");
741 battr->set_type(onnx::AttributeProto::GRAPH);
742 *battr->mutable_g() = body;
745 casadi_int si = 0, ci = 0;
746 for (casadi_int j = 0; j < base.n_out(); ++j) {
748 work_to_onnx[o[j]] = state_out[si++];
750 std::string out = out_prefix + std::to_string(j);
751 emit_reshape(container, scan_out[ci++], {n * base.size2_out(j), base.size1_out(j)}, out,
752 out_prefix +
"_outsh" + std::to_string(j));
753 work_to_onnx[o[j]] = out;
758 onnx::GraphProto Onnx::build_if_branch(
const Function& f,
759 const std::vector<std::string>& arg_names,
const std::string& prefix) {
762 g.set_name(f.name() +
"_branch");
763 std::map<casadi_int, std::string> work_to_onnx;
765 casadi_int n_instr = f.n_instructions();
766 for (casadi_int k = 0; k < n_instr; ++k) {
767 casadi_int op = f.instruction_id(k);
768 std::vector<casadi_int> o = f.instruction_output(k);
769 std::vector<casadi_int> i_vec = f.instruction_input(k);
770 std::string node_output =
"n" + std::to_string(k);
773 work_to_onnx[o[0]] = arg_names.at(i_vec[0]);
777 std::string out_name =
"branch_out_" + std::to_string(o[0]);
778 MX dep = f.instruction_MX(k).dep(0);
779 Sparsity out_sp = f.sparsity_out(o[0]);
780 emit_output_node(&g, work_to_onnx[i_vec[0]], dep.size1(), dep.size2(), out_sp,
781 out_name,
"branch_out_rs_" + std::to_string(k));
782 onnx::ValueInfoProto* vi = g.add_output();
783 vi->set_name(out_name);
784 set_real_tensor_type(vi, out_sp);
788 Function called = f.instruction_MX(k).which_function();
789 if (is_map_function(called)) {
790 export_map(&g, called, i_vec, o, work_to_onnx, node_output +
"_out");
791 }
else if (is_reduce_map_function(called)) {
792 export_reduce_map(&g, called, i_vec, o, work_to_onnx, node_output +
"_out");
793 }
else if (is_if_else_function(called)) {
794 export_if(&g, called, i_vec, o, work_to_onnx, node_output +
"_out");
796 assert_not_control_flow(called);
797 export_call(&g, called, i_vec, o, work_to_onnx, node_output +
"_out");
801 if (process_operation(&g, f, op, k, i_vec, o, work_to_onnx, node_output))
continue;
803 casadi_error(
"ONNX export: unsupported operation code " + std::to_string(op) +
804 " in if_else branch of '" + f.name() +
"'");
810 template<
typename Container>
811 void Onnx::export_if(Container* container,
const Function& switch_fn,
812 const std::vector<casadi_int>& i_vec,
813 const std::vector<casadi_int>& o,
814 std::map<casadi_int, std::string>& work_to_onnx,
815 const std::string& out_prefix) {
818 Dict info = switch_fn.info();
819 std::vector<Function> cases = info.at(
"f");
820 Function f_then = info.at(
"f_def");
821 casadi_assert(cases.size() == 1,
822 "ONNX export: only 2-way if_else is supported (got a multi-case Switch).");
823 Function f_else = cases[0];
828 std::vector<std::string> arg_names;
829 for (
size_t j = 1; j < i_vec.size(); ++j) {
830 std::string alias = out_prefix +
"_arg" + std::to_string(j - 1);
831 onnx::NodeProto*
id = container->add_node();
832 id->set_op_type(
"Identity");
833 id->add_input(work_to_onnx[i_vec[j]]);
834 id->add_output(alias);
835 arg_names.push_back(alias);
838 std::string cond = out_prefix +
"_cond";
839 onnx::NodeProto* cast = container->add_node();
840 cast->set_op_type(
"Cast");
841 cast->add_input(work_to_onnx[i_vec[0]]);
842 cast->add_output(cond);
845 onnx::NodeProto* if_node = container->add_node();
846 if_node->set_op_type(
"If");
847 if_node->add_input(cond);
848 for (
size_t j = 0; j < o.size(); ++j) {
849 std::string out = out_prefix + std::to_string(j);
850 if_node->add_output(out);
851 work_to_onnx[o[j]] = out;
853 onnx::AttributeProto* then_attr = if_node->add_attribute();
854 then_attr->set_name(
"then_branch");
855 then_attr->set_type(onnx::AttributeProto::GRAPH);
856 *then_attr->mutable_g() = build_if_branch(f_then, arg_names, out_prefix +
"_t_");
858 onnx::AttributeProto* else_attr = if_node->add_attribute();
859 else_attr->set_name(
"else_branch");
860 else_attr->set_type(onnx::AttributeProto::GRAPH);
861 *else_attr->mutable_g() = build_if_branch(f_else, arg_names, out_prefix +
"_e_");
864 template<
typename Container>
865 void Onnx::export_call(Container* container,
const Function& called_func,
866 const std::vector<casadi_int>& i_vec,
867 const std::vector<casadi_int>& o,
868 std::map<casadi_int, std::string>& work_to_onnx,
869 const std::string& out_prefix) {
870 std::string func_name = called_func.name();
871 std::string domain =
"casadi";
875 onnx::FunctionProto* func_proto = function_to_function_proto(called_func, domain);
876 *
model_.add_functions() = *func_proto;
881 onnx::NodeProto* call_node = container->add_node();
882 call_node->set_op_type(func_name);
883 call_node->set_domain(domain);
884 for (casadi_int idx : i_vec) call_node->add_input(work_to_onnx[idx]);
885 for (
size_t j = 0; j < o.size(); ++j) {
886 std::string out = out_prefix + std::to_string(j);
887 call_node->add_output(out);
888 work_to_onnx[o[j]] = out;
892 onnx::FunctionProto* Onnx::function_to_function_proto(
894 const std::string& domain) {
897 std::unique_ptr<onnx::FunctionProto> func_guard(
new onnx::FunctionProto());
898 onnx::FunctionProto* func = func_guard.get();
899 func->set_name(f.name());
900 func->set_domain(domain);
904 onnx::OperatorSetIdProto* opset = func->add_opset_import();
905 opset->set_domain(
"");
906 opset->set_version(16);
908 onnx::OperatorSetIdProto* casadi_opset = func->add_opset_import();
909 casadi_opset->set_domain(
"casadi");
910 casadi_opset->set_version(1);
913 uout() <<
" Converting Function '" << f.name() <<
"' to ONNX FunctionProto"
914 <<
" (domain: " << domain <<
")" << std::endl;
915 uout() <<
" Inputs: " << f.n_in() <<
", Outputs: " << f.n_out()
916 <<
", Instructions: " << f.n_instructions() << std::endl;
920 for (casadi_int i = 0; i < f.n_in(); ++i) {
925 for (casadi_int i = 0; i < f.n_out(); ++i) {
930 std::map<casadi_int, std::string> work_to_onnx;
935 std::set<casadi_int> written_outputs;
938 casadi_int n_instr = f.n_instructions();
939 for (casadi_int k = 0; k < n_instr; ++k) {
940 casadi_int op = f.instruction_id(k);
941 std::vector<casadi_int> o = f.instruction_output(k);
942 std::vector<casadi_int> i_vec = f.instruction_input(k);
944 std::string node_output =
"n" + std::to_string(k);
948 casadi_assert(f.instruction_MX(k).numel() == f.numel_in(i_vec[0]),
949 "ONNX export: nested functions with multi-segment inputs (e.g. mapaccum) "
950 "are not supported.");
958 casadi_assert(written_outputs.insert(o[0]).second,
959 "ONNX export: nested functions with multi-segment outputs (e.g. mapaccum) "
960 "are not supported.");
961 MX dep = f.instruction_MX(k).dep(0);
962 Sparsity out_sp = f.sparsity_out(o[0]);
963 emit_output_node(func, work_to_onnx[i_vec[0]], dep.size1(), dep.size2(), out_sp,
970 Function called_func = f.instruction_MX(k).which_function();
971 std::string out_prefix =
"call_" + called_func.name() +
"_" + std::to_string(k) +
"_out";
972 if (is_map_function(called_func)) {
973 export_map(func, called_func, i_vec, o, work_to_onnx, out_prefix);
974 }
else if (is_reduce_map_function(called_func)) {
975 export_reduce_map(func, called_func, i_vec, o, work_to_onnx, out_prefix);
976 }
else if (is_if_else_function(called_func)) {
977 export_if(func, called_func, i_vec, o, work_to_onnx, out_prefix);
979 assert_not_control_flow(called_func);
980 export_call(func, called_func, i_vec, o, work_to_onnx, out_prefix);
986 if (process_operation([&]() {
return func->add_node(); },
987 f, op, k, i_vec, o, work_to_onnx, node_output)) {
992 casadi_error(
"ONNX export: Unsupported operation code " + std::to_string(op) +
993 " in function '" + f.name() +
"' at instruction " + std::to_string(k));
997 uout() <<
" Created " << func->node_size() <<
" nodes in FunctionProto" << std::endl;
1000 return func_guard.release();
casadi_int n_instructions() const
Number of instruction in the algorithm (SXFunction/MXFunction)
casadi_int size2_out(casadi_int ind) const
Get output dimension.
const Sparsity & sparsity_out(casadi_int ind) const
Get sparsity of a given output.
const std::vector< std::string > & name_in() const
Get input scheme.
const std::string & name() const
Name of the function.
std::vector< casadi_int > instruction_input(casadi_int k) const
Locations in the work vector for the inputs of the instruction.
const Sparsity & sparsity_in(casadi_int ind) const
Get sparsity of a given input.
std::vector< casadi_int > instruction_output(casadi_int k) const
Location in the work vector for the output of the instruction.
MX instruction_MX(casadi_int k) const
Get the MX node corresponding to an instruction (MXFunction)
casadi_int n_out() const
Get the number of function outputs.
casadi_int n_in() const
Get the number of function inputs.
casadi_int size1_out(casadi_int ind) const
Get output dimension.
casadi_int instruction_id(casadi_int k) const
Identifier index of the instruction (SXFunction/MXFunction)
const std::vector< std::string > & name_out() const
Get output scheme.
casadi_int size2() const
Get the second dimension (i.e. number of columns)
casadi_int size1() const
Get the first dimension (i.e. number of rows)
bool verbose_
Verbose – for debugging.
const Sparsity & sparsity() const
Get the sparsity pattern.
Function which_function() const
Get function - only valid when is_call() is true.
MX dep(casadi_int ch=0) const
Get the nth dependency as MX.
std::set< std::string > exported_functions_
Track which functions have been exported as FunctionProto.
std::string class_name() const override
Readable name of the internal class.
onnx::ModelProto model_
ONNX model protocol buffer.
bool has_model_
Whether a model has been loaded.
void load(const Function &f)
Load a CasADi Function and convert to the ONNX representation.
std::string class_name() const
Get class name.
casadi_int size1() const
Get the number of rows.
static Sparsity dense(casadi_int nrow, casadi_int ncol=1)
Create a dense rectangular sparsity pattern *.
casadi_int size2() const
Get the number of columns.
bool is_dense() const
Is dense?
onnx::NodeProto * create_unary_node(AddNodeFn add_node, const std::string &op_type, const std::string &input, const std::string &output)
Create unary operation ONNX node (callback-based)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
std::string onnx_output_name(const Function &f, casadi_int i)
onnx::NodeProto * create_binary_node(AddNodeFn add_node, const std::string &op_type, const std::string &input1, const std::string &input2, const std::string &output)
Create binary operation ONNX node (callback-based)
void add_int_attribute(onnx::NodeProto *node, const std::string &name, casadi_int value)
Add an integer attribute (e.g. axis) to a node.
static void prefix_graph_names(onnx::GraphProto *g, const std::string &prefix)
std::function< onnx::NodeProto *()> AddNodeFn
Callback type for adding nodes to a container (GraphProto or FunctionProto)
std::string onnx_input_name(const Function &f, casadi_int i)
Function input/output name, or a generated fallback when unnamed.
void add_ints_attribute(onnx::NodeProto *node, const std::string &name, const std::vector< casadi_int > &values)
Add an integer-list attribute (e.g. scan_input_axes) to a node.