26 #include "onnx_model.hpp"
76 for (
const auto& entry :
op_map) {
77 if (entry.casadi_op == op)
return &entry;
84 for (
const auto& entry :
op_map) {
85 if (entry.onnx_name == onnx_name)
return &entry;
96 const std::string& output_name,
97 onnx::TensorProto::DataType data_type) {
98 onnx::NodeProto* node = add_node();
99 node->set_op_type(
"Constant");
100 node->add_output(output_name);
101 onnx::AttributeProto* attr = node->add_attribute();
102 attr->set_name(
"value");
103 attr->set_type(onnx::AttributeProto::TENSOR);
104 onnx::TensorProto* tensor = attr->mutable_t();
105 tensor->set_data_type(data_type);
111 const std::vector<casadi_int>& data,
112 std::vector<casadi_int> dims = {}) {
114 if (dims.empty()) dims = {
static_cast<casadi_int
>(data.size())};
115 for (casadi_int d : dims) t->add_dims(d);
116 for (casadi_int v : data) t->add_int64_data(v);
120 void Onnx::add_real_constant(
AddNodeFn add_node,
const std::string& name,
121 const std::vector<double>& data,
122 const std::vector<casadi_int>& dims) {
124 for (casadi_int d : dims) t->add_dims(d);
125 if (real_type() == onnx::TensorProto::FLOAT) {
126 for (
double v : data) t->add_float_data(
static_cast<float>(v));
128 for (
double v : data) t->add_double_data(v);
133 void Onnx::add_sparse_constant(
AddNodeFn add_node,
const std::string& name,
const DM& dm) {
134 onnx::NodeProto* node = add_node();
135 node->set_op_type(
"Constant");
136 node->add_output(name);
137 onnx::AttributeProto* attr = node->add_attribute();
138 attr->set_name(
"sparse_value");
139 attr->set_type(onnx::AttributeProto::SPARSE_TENSOR);
140 fill_sparse_tensor(attr->mutable_sparse_tensor(), name, dm);
143 void Onnx::fill_sparse_tensor(onnx::SparseTensorProto* st,
const std::string& name,
144 const DM& dm)
const {
148 st->add_dims(dm.size2());
149 st->add_dims(dm.size1());
150 casadi_int nnz = dm.nnz();
151 std::vector<casadi_int> row = dm.sparsity().get_row();
152 std::vector<casadi_int> col = dm.sparsity().get_col();
153 const std::vector<double>& vals = dm.nonzeros();
155 onnx::TensorProto* vt = st->mutable_values();
157 vt->set_data_type(real_type());
159 if (real_type() == onnx::TensorProto::FLOAT) {
160 for (
double v : vals) vt->add_float_data(
static_cast<float>(v));
162 for (
double v : vals) vt->add_double_data(v);
165 onnx::TensorProto* it = st->mutable_indices();
166 it->set_data_type(onnx::TensorProto::INT64);
169 for (casadi_int k = 0; k < nnz; ++k) {
170 it->add_int64_data(col[k]);
171 it->add_int64_data(row[k]);
177 onnx::AttributeProto* attr = node->add_attribute();
178 attr->set_name(name);
179 attr->set_type(onnx::AttributeProto::INT);
185 const std::vector<casadi_int>& values) {
186 onnx::AttributeProto* attr = node->add_attribute();
187 attr->set_name(name);
188 attr->set_type(onnx::AttributeProto::INTS);
189 for (casadi_int v : values) attr->add_ints(v);
194 const std::vector<casadi_int>& data, std::vector<casadi_int> dims) {
195 add_int_constant([graph]() {
return graph->add_node(); }, name, data, dims);
200 const std::string& indices,
const std::string& output,
201 casadi_int axis = 0) {
202 onnx::NodeProto* node = add_node();
203 node->set_op_type(
"Gather");
204 node->add_input(data);
205 node->add_input(indices);
206 node->add_output(output);
213 const std::string& data,
const std::vector<casadi_int>& starts,
214 const std::vector<casadi_int>& ends,
const std::vector<casadi_int>& axes,
215 const std::vector<casadi_int>& steps,
const std::string& output) {
221 onnx::NodeProto* node = add_node();
222 node->set_op_type(
"Slice");
223 node->add_input(data);
224 node->add_input(uniq +
"_starts");
225 node->add_input(uniq +
"_ends");
226 node->add_input(uniq +
"_axes");
227 if (!steps.empty()) node->add_input(uniq +
"_steps");
228 node->add_output(output);
237 const std::vector<casadi_int>& dims,
const std::string& output,
238 const std::string& uniq) {
239 std::vector<casadi_int> rev(dims.rbegin(), dims.rend());
247 const std::string& op_type,
248 const std::string& input1,
249 const std::string& input2,
250 const std::string& output) {
251 onnx::NodeProto* node = add_node();
252 node->set_op_type(op_type);
253 node->add_input(input1);
254 node->add_input(input2);
255 node->add_output(output);
261 const std::string& op_type,
262 const std::string& input,
263 const std::string& output) {
264 onnx::NodeProto* node = add_node();
265 node->set_op_type(op_type);
266 node->add_input(input);
267 node->add_output(output);
274 const std::string& input,
275 const std::string& output,
276 onnx::TensorProto::DataType to_type) {
277 onnx::NodeProto* node = add_node();
278 node->set_op_type(
"Cast");
279 node->add_input(input);
280 node->add_output(output);
281 onnx::AttributeProto* attr = node->add_attribute();
282 attr->set_name(
"to");
283 attr->set_type(onnx::AttributeProto::INT);
284 attr->set_i(to_type);
294 const std::string& A,
const std::string& B,
const std::string& C,
295 const std::string& output,
bool transA =
false,
bool transB =
false) {
296 onnx::NodeProto* node = add_node();
297 node->set_op_type(
"Gemm");
300 if (!
C.empty()) node->add_input(
C);
301 node->add_output(output);
310 const std::string& cond,
311 const std::string& if_true,
312 const std::string& if_false,
313 const std::string& output) {
314 onnx::NodeProto* node = add_node();
315 node->set_op_type(
"Where");
316 node->add_input(cond);
317 node->add_input(if_true);
318 node->add_input(if_false);
319 node->add_output(output);
325 onnx::GraphProto* graph,
326 const std::string& op_type,
327 const std::string& input1,
328 const std::string& input2,
329 const std::string& output) {
331 op_type, input1, input2, output);
335 onnx::GraphProto* graph,
336 const std::string& op_type,
337 const std::string& input,
338 const std::string& output) {
340 op_type, input, output);
346 void Onnx::emit_nonzero_remap(
AddNodeFn add_node,
const std::string& data,
347 const Sparsity& sp_in,
const Sparsity& out_sp,
348 std::vector<casadi_int> idx,
const std::string& uniq,
349 const std::string& node_output) {
350 casadi_int numel_in = sp_in.size1() * sp_in.size2();
351 std::vector<casadi_int> loc = sp_in.find();
352 bool has_fill =
false;
353 for (casadi_int& v : idx) {
354 if (v < 0) { v = numel_in; has_fill =
true; }
360 std::string flat = data;
361 if (sp_in.size2() != 1) {
362 flat =
"gnz_flat_" + uniq;
367 std::string zc =
"gnz_fz_" + uniq, fl2 =
"gnz_flz_" + uniq;
368 add_real_constant(add_node, zc, {0.0}, {1, 1});
369 onnx::NodeProto* cc = add_node();
370 cc->set_op_type(
"Concat");
371 cc->add_input(flat); cc->add_input(zc); cc->add_output(fl2);
377 casadi_int numel_out = out_sp.size1() * out_sp.size2();
378 bool dense_out = (
static_cast<casadi_int
>(idx.size()) == numel_out);
380 std::string gathered = (dense_out && out_sp.size2() == 1) ? node_output : (
"gnz_g_" + uniq);
385 if (out_sp.size2() != 1) {
394 std::vector<casadi_int> ofind = out_sp.find();
395 std::string zname =
"gnz_zeros_" + uniq, fname =
"gnz_ofind_" + uniq;
396 std::string scat =
"gnz_scat_" + uniq;
397 add_real_constant(add_node, zname, std::vector<double>(numel_out, 0.0), {1, numel_out});
398 add_int_constant(add_node, fname, ofind, {1,
static_cast<casadi_int
>(ofind.size())});
399 onnx::NodeProto* sc = add_node();
400 sc->set_op_type(
"ScatterElements");
401 sc->add_input(zname); sc->add_input(fname); sc->add_input(gathered);
402 sc->add_output(scat);
404 std::string dframe =
"gnz_dense_" + uniq;
407 emit_sparsity_restore(add_node, dframe,
Sparsity::dense(out_sp.size1(),
408 out_sp.size2()), out_sp,
"gnzr_" + uniq, node_output);
413 std::string Onnx::emit_sparsity_restore(
AddNodeFn add_node,
const std::string& value,
414 const Sparsity& value_sp,
const Sparsity& target_sp,
415 const std::string& uniq,
416 const std::string& final_output) {
418 if (value_sp == target_sp)
return value;
420 if (target_sp.is_dense()) {
423 std::string zc =
"sr_dz_" + uniq;
424 add_real_constant(add_node, zc,
425 std::vector<double>(target_sp.size1() * target_sp.size2(), 0.0),
426 {target_sp.size2(), target_sp.size1()});
439 bool subset = value_sp.is_subset(target_sp);
440 bool superset = target_sp.is_subset(value_sp);
442 std::string zeros =
"sr_zeros_" + uniq;
443 add_sparse_constant(add_node, zeros,
DM(target_sp, 0.0));
448 std::string ones =
"sr_ones_" + uniq;
449 add_sparse_constant(add_node, ones,
DM(target_sp, 1.0));
454 std::string ones =
"sr_ones_" + uniq;
455 add_sparse_constant(add_node, ones,
DM(target_sp, 1.0));
456 std::string mul =
"sr_mul_" + uniq;
458 std::string zeros =
"sr_zeros_" + uniq;
459 add_sparse_constant(add_node, zeros,
DM(target_sp, 0.0));
465 void Onnx::emit_blockdiag(
AddNodeFn add_node,
const std::vector<std::string>& names,
466 const std::vector<casadi_int>& row_off,
467 const std::vector<casadi_int>& col_off,
468 const std::vector<casadi_int>& br,
const std::vector<casadi_int>& bc,
469 casadi_int R, casadi_int C,
const std::string& output,
470 const std::string& uniq) {
471 std::vector<std::string> padded;
472 for (casadi_int b = 0; b < static_cast<casadi_int>(names.size()); ++b) {
475 std::string pn = uniq +
"_pp" + std::to_string(b);
477 C - col_off[b] - bc[b], R - row_off[b] - br[b]});
478 std::string pd = uniq +
"_pad" + std::to_string(b);
479 onnx::NodeProto* pad = add_node();
480 pad->set_op_type(
"Pad");
481 pad->add_input(names[b]);
484 padded.push_back(pd);
486 if (padded.size() == 1) {
489 onnx::NodeProto*
sum = add_node();
490 sum->set_op_type(
"Sum");
491 for (
const auto& p : padded)
sum->add_input(p);
492 sum->add_output(output);
497 bool Onnx::process_operation(
502 const std::vector<casadi_int>& i_vec,
503 const std::vector<casadi_int>& o_vec,
504 std::map<casadi_int, std::string>& work_to_onnx,
505 const std::string& node_output) {
507 onnx::NodeProto* node =
nullptr;
512 if (mapping->arity == 1 && i_vec.size() >= 1 && o_vec.size() == 1) {
513 create_unary_node(add_node, mapping->onnx_name, work_to_onnx[i_vec[0]], node_output);
514 work_to_onnx[o_vec[0]] = node_output;
516 }
else if (mapping->arity == 2 && i_vec.size() >= 2 && o_vec.size() == 1) {
518 work_to_onnx[i_vec[1]], node_output);
519 work_to_onnx[o_vec[0]] = node_output;
530 MX mx_input = f.instruction_MX(k);
531 Dict info = mx_input.info();
532 casadi_int offset = info[
"offset"];
533 casadi_int input_numel = f.numel_in(i_vec[0]);
534 casadi_int output_numel = mx_input.numel();
536 if (output_numel < input_numel) {
539 std::string uniq = std::to_string(k);
540 std::string flat = input_name;
541 if (f.size2_in(i_vec[0]) != 1) {
542 flat =
"in_flat_" + uniq;
545 std::vector<casadi_int> idx;
546 for (casadi_int t = 0; t < output_numel; ++t) idx.push_back(offset + t);
547 std::string index_name =
"input_idx_" + uniq;
549 bool out_is_col = (mx_input.size2() == 1);
550 std::string gathered = out_is_col ? node_output : (
"in_g_" + uniq);
554 node_output,
"in_out_" + uniq);
561 work_to_onnx[o_vec[0]] = input_name;
564 work_to_onnx[o_vec[0]] = node_output;
576 DM dm_const =
static_cast<DM>(f.instruction_MX(k));
577 if (dm_const.is_dense()) {
579 std::vector<double> data(dm_const->begin(), dm_const->end());
580 add_real_constant(add_node, node_output, data, {dm_const.size2(), dm_const.size1()});
583 add_sparse_constant(add_node, node_output, dm_const);
585 work_to_onnx[o_vec[0]] = node_output;
592 work_to_onnx[i_vec[0]], node_output);
593 work_to_onnx[o_vec[0]] = node_output;
600 work_to_onnx[i_vec[0]], node_output);
601 work_to_onnx[o_vec[0]] = node_output;
607 std::string const_name =
"const_2_" + std::to_string(k);
608 add_real_constant(add_node, const_name, {2.0});
610 work_to_onnx[o_vec[0]] = node_output;
616 work_to_onnx[o_vec[0]] = node_output;
620 create_unary_node(add_node,
"ReduceLogSumExp", work_to_onnx[i_vec[0]], node_output);
621 work_to_onnx[o_vec[0]] = node_output;
626 std::string one =
"one_" + std::to_string(k), s =
"log1p_" + std::to_string(k);
627 add_real_constant(add_node, one, {1.0});
630 work_to_onnx[o_vec[0]] = node_output;
636 std::string e =
"exp_" + std::to_string(k), one =
"one_" + std::to_string(k);
638 add_real_constant(add_node, one, {1.0});
640 work_to_onnx[o_vec[0]] = node_output;
646 std::string x2 =
"hx_" + std::to_string(k), y2 =
"hy_" + std::to_string(k);
647 std::string s =
"hs_" + std::to_string(k);
648 create_binary_node(add_node,
"Mul", work_to_onnx[i_vec[0]], work_to_onnx[i_vec[0]], x2);
649 create_binary_node(add_node,
"Mul", work_to_onnx[i_vec[1]], work_to_onnx[i_vec[1]], y2);
652 work_to_onnx[o_vec[0]] = node_output;
662 work_to_onnx[o_vec[0]] = node_output;
667 MX mx_solve = f.instruction_MX(k);
668 Dict info = mx_solve.info();
669 bool tr = info[
"tr"];
670 casadi_error(
"ONNX export: OP_SOLVE (linear solver) is not supported. "
671 "ONNX does not provide native linear algebra solvers. "
672 "Consider using explicit matrix operations or iterative methods. "
673 "Transpose flag was: " + std::string(tr ?
"true" :
"false"));
679 std::string abs_result =
"abs_" + std::to_string(k);
682 work_to_onnx[o_vec[0]] = node_output;
688 std::string b =
"cmp_" + std::to_string(k);
689 create_binary_node(add_node,
"Less", work_to_onnx[i_vec[0]], work_to_onnx[i_vec[1]], b);
691 work_to_onnx[o_vec[0]] = node_output;
696 std::string b =
"cmp_" + std::to_string(k);
698 work_to_onnx[i_vec[1]], b);
700 work_to_onnx[o_vec[0]] = node_output;
705 std::string b =
"cmp_" + std::to_string(k);
706 create_binary_node(add_node,
"Equal", work_to_onnx[i_vec[0]], work_to_onnx[i_vec[1]], b);
708 work_to_onnx[o_vec[0]] = node_output;
714 std::string b =
"bin_" + std::to_string(k);
715 create_cast_node(add_node, work_to_onnx[i_vec[0]], b, onnx::TensorProto::BOOL);
716 std::string r =
"bres_" + std::to_string(k);
719 work_to_onnx[o_vec[0]] = node_output;
725 std::string b0 =
"bin0_" + std::to_string(k), b1 =
"bin1_" + std::to_string(k);
726 create_cast_node(add_node, work_to_onnx[i_vec[0]], b0, onnx::TensorProto::BOOL);
727 create_cast_node(add_node, work_to_onnx[i_vec[1]], b1, onnx::TensorProto::BOOL);
728 std::string r =
"bres_" + std::to_string(k);
731 work_to_onnx[o_vec[0]] = node_output;
738 work_to_onnx[o_vec[0]] = node_output;
743 node->set_op_type(
"Mod");
744 node->add_input(work_to_onnx[i_vec[0]]);
745 node->add_input(work_to_onnx[i_vec[1]]);
746 node->add_output(node_output);
748 work_to_onnx[o_vec[0]] = node_output;
754 std::string sign_result =
"sign_" + std::to_string(k);
755 std::string abs_result =
"abs_" + std::to_string(k);
759 work_to_onnx[o_vec[0]] = node_output;
766 work_to_onnx[o_vec[0]] = node_output;
771 std::string equal_result =
"equal_" + std::to_string(k);
774 std::string r =
"bres_" + std::to_string(k);
777 work_to_onnx[o_vec[0]] = node_output;
783 std::string cond =
"cond_" + std::to_string(k);
784 create_cast_node(add_node, work_to_onnx[i_vec[0]], cond, onnx::TensorProto::BOOL);
785 std::string zero_name =
"const_0_" + std::to_string(k);
786 add_real_constant(add_node, zero_name, {0.0});
787 create_where_node(add_node, cond, work_to_onnx[i_vec[1]], zero_name, node_output);
788 work_to_onnx[o_vec[0]] = node_output;
794 std::string mul_result =
"mul_" + std::to_string(k);
798 work_to_onnx[o_vec[0]] = node_output;
805 std::string xa =
"bilin_" + std::to_string(k);
806 create_gemm_node(add_node, work_to_onnx[i_vec[1]], work_to_onnx[i_vec[0]],
"", xa,
true);
809 work_to_onnx[o_vec[0]] = node_output;
816 std::string xyt =
"rank1_" + std::to_string(k);
817 create_gemm_node(add_node, work_to_onnx[i_vec[2]], work_to_onnx[i_vec[3]],
"", xyt,
819 std::string scaled =
"rank1s_" + std::to_string(k);
822 work_to_onnx[o_vec[0]] = node_output;
830 MX mx_e = f.instruction_MX(k);
831 Dict info = mx_e.info();
832 std::vector<casadi_int> la = info[
"a"], lb = info[
"b"], lc = info[
"c"];
833 std::vector<casadi_int> da = info[
"dim_a"], db = info[
"dim_b"], dc = info[
"dim_c"];
834 casadi_assert(da.size() <= 2 && db.size() <= 2 && dc.size() <= 2,
835 "ONNX export: einstein with >2-index operands is not supported.");
838 std::map<casadi_int, char> letter;
840 for (
const std::vector<casadi_int>& labs : {la, lb, lc}) {
841 for (casadi_int L : labs)
if (!letter.count(L)) letter[L] = nxt++;
843 std::string sa, sb, sc;
844 for (casadi_int L : la) sa += letter[L];
845 for (casadi_int L : lb) sb += letter[L];
846 for (casadi_int L : lc) sc += letter[L];
851 std::reverse(sa.begin(), sa.end());
852 std::reverse(sb.begin(), sb.end());
853 std::reverse(sc.begin(), sc.end());
855 std::string a_re =
"ein_a_" + std::to_string(k), b_re =
"ein_b_" + std::to_string(k);
859 std::string ein_out =
"ein_o_" + std::to_string(k);
861 node->set_op_type(
"Einsum");
862 node->add_input(a_re);
863 node->add_input(b_re);
864 node->add_output(ein_out);
865 onnx::AttributeProto* eq = node->add_attribute();
866 eq->set_name(
"equation");
867 eq->set_type(onnx::AttributeProto::STRING);
868 eq->set_s(sa +
"," + sb +
"->" + sc);
871 casadi_int prod_c = 1;
872 for (casadi_int d : dc) prod_c *= d;
873 std::string ein_vec =
"ein_v_" + std::to_string(k);
876 work_to_onnx[o_vec[0]] = node_output;
893 MX mx_kron = f.instruction_MX(k);
894 casadi_int ra = mx_kron.dep(0).size1(), ca = mx_kron.dep(0).size2();
895 casadi_int rb = mx_kron.dep(1).size1(), cb = mx_kron.dep(1).size2();
896 std::string uniq = std::to_string(k);
897 std::string tag =
"kron" + uniq;
902 auto emit_reshape = [&](
const std::string& data,
const std::vector<casadi_int>& shape,
903 const std::string& out,
const std::string& role) ->
void {
906 rn->add_input(out +
"_s");
907 rn->set_name(tag +
"_" + role);
910 std::string a4 = tag +
"_A4", b4 = tag +
"_B4", p = tag +
"_P";
911 emit_reshape(work_to_onnx[i_vec[0]], {ca, 1, ra, 1}, a4,
"A4");
912 emit_reshape(work_to_onnx[i_vec[1]], {1, cb, 1, rb}, b4,
"B4");
914 mul->set_name(tag +
"_P");
915 emit_reshape(p, {ca * cb, ra * rb}, node_output,
"R");
917 work_to_onnx[o_vec[0]] = node_output;
927 auto sp = f.instruction_MX(k).sparsity();
930 node_output,
"rs_" + std::to_string(k));
932 std::string rs =
"rs_d_" + std::to_string(k);
934 rs,
"rs_" + std::to_string(k));
935 std::string r = emit_sparsity_restore(add_node, rs,
937 "rsr_" + std::to_string(k), node_output);
938 work_to_onnx[o_vec[0]] = r;
941 work_to_onnx[o_vec[0]] = node_output;
950 node->set_op_type(
"Concat");
951 for (casadi_int idx : i_vec) node->add_input(work_to_onnx[idx]);
952 node->add_output(node_output);
954 work_to_onnx[o_vec[0]] = node_output;
962 MX mx_dc = f.instruction_MX(k);
963 std::string uniq = std::to_string(k);
964 std::vector<std::string> names;
965 std::vector<casadi_int> row_off, col_off, brs, bcs;
966 casadi_int ro = 0, co = 0;
967 for (casadi_int j = 0; j < static_cast<casadi_int>(i_vec.size()); ++j) {
968 casadi_int br = mx_dc.dep(j).size1(), bc = mx_dc.dep(j).size2();
969 names.push_back(work_to_onnx[i_vec[j]]);
970 row_off.push_back(ro); col_off.push_back(co); brs.push_back(br); bcs.push_back(bc);
973 Sparsity dc_sp = mx_dc.sparsity();
974 if (dc_sp.is_dense()) {
975 emit_blockdiag(add_node, names, row_off, col_off, brs, bcs,
976 mx_dc.size1(), mx_dc.size2(), node_output,
"dc" + uniq);
979 std::string tmp =
"dc_pre_" + uniq;
980 emit_blockdiag(add_node, names, row_off, col_off, brs, bcs,
981 mx_dc.size1(), mx_dc.size2(), tmp,
"dc" + uniq);
982 std::string r = emit_sparsity_restore(add_node, tmp,
985 work_to_onnx[o_vec[0]] = r;
988 work_to_onnx[o_vec[0]] = node_output;
1000 MX mx_split = f.instruction_MX(k);
1002 Function split_out = mx_split.info()[
"output"];
1004 std::vector<casadi_int> split_sizes;
1005 for (casadi_int j = 0; j < split_out.n_out(); ++j) {
1006 split_sizes.push_back(horz ? split_out.size2_out(j) : split_out.size1_out(j));
1009 std::string split_sizes_name =
"split_sizes_" + std::to_string(k);
1013 node->set_op_type(
"Split");
1014 node->add_input(work_to_onnx[i_vec[0]]);
1015 node->add_input(split_sizes_name);
1019 for (casadi_int j = 0; j < o_vec.size(); ++j) {
1020 std::string output_name =
"n" + std::to_string(k) +
"_out" + std::to_string(j);
1021 node->add_output(output_name);
1022 work_to_onnx[o_vec[j]] = output_name;
1029 MX mx_repmat = f.instruction_MX(k);
1030 casadi_int n = mx_repmat.size2() / mx_repmat.dep(0).size2();
1032 std::string repeats_name =
"repeats_" + std::to_string(k);
1036 node->set_op_type(
"Tile");
1037 node->add_input(work_to_onnx[i_vec[0]]);
1038 node->add_input(repeats_name);
1039 node->add_output(node_output);
1040 work_to_onnx[o_vec[0]] = node_output;
1046 MX mx_repsum = f.instruction_MX(k);
1047 casadi_int output_cols = mx_repsum.size2();
1048 casadi_int n = mx_repsum.dep(0).size2() / output_cols;
1050 std::string split_sizes_name =
"split_sizes_" + std::to_string(k);
1051 add_int_constant(add_node, split_sizes_name, std::vector<casadi_int>(n, output_cols));
1054 node->set_op_type(
"Split");
1055 node->add_input(work_to_onnx[i_vec[0]]);
1056 node->add_input(split_sizes_name);
1060 onnx::NodeProto* sum_node = add_node();
1061 sum_node->set_op_type(
"Sum");
1062 for (casadi_int i = 0; i < n; ++i) {
1063 std::string out_name =
"split_" + std::to_string(k) +
"_" + std::to_string(i);
1064 node->add_output(out_name);
1065 sum_node->add_input(out_name);
1067 sum_node->add_output(node_output);
1068 work_to_onnx[o_vec[0]] = node_output;
1077 MX mx_getnonzeros = f.instruction_MX(k);
1078 Dict info = mx_getnonzeros.info();
1079 std::string data = work_to_onnx[i_vec[0]];
1080 std::string uniq = std::to_string(k);
1083 std::vector<casadi_int> idx;
1084 if (info.count(
"nz")) {
1086 }
else if (info.count(
"slice")) {
1087 Dict s = info[
"slice"];
1088 casadi_int start = s[
"start"], stop = s[
"stop"], step = s[
"step"];
1089 for (casadi_int i = start; i < stop; i += step) idx.push_back(i);
1090 }
else if (info.count(
"inner")) {
1091 Dict inner_info = info[
"inner"], outer_info = info[
"outer"];
1092 casadi_int is = inner_info[
"start"], ip = inner_info[
"stop"], ist = inner_info[
"step"];
1093 casadi_int os = outer_info[
"start"], op2 = outer_info[
"stop"], ost = outer_info[
"step"];
1094 for (casadi_int o = os; o < op2; o += ost) {
1095 for (casadi_int i = is; i < ip; i += ist) idx.push_back(o + i);
1108 const Sparsity& sp_in = mx_getnonzeros.dep(0).sparsity();
1109 casadi_int m_in = sp_in.size1(), n_in = sp_in.size2();
1110 if (sp_in.is_dense() && !idx.empty()) {
1113 casadi_int N =
static_cast<casadi_int
>(idx.size());
1114 casadi_int c0 = idx[0] / m_in, r0 = idx[0] % m_in;
1116 while (k < N && idx[k] / m_in == c0) ++k;
1117 bool grid = (N % k == 0);
1118 casadi_int nc = grid ? N / k : 0;
1119 casadi_int rstep = (k > 1) ? (idx[1] % m_in) - r0 : 1;
1120 casadi_int cstep = (grid && nc > 1) ? idx[k] / m_in - c0 : 1;
1121 if (grid && rstep > 0 && cstep > 0) {
1122 for (casadi_int cj = 0; cj < nc && grid; ++cj)
1123 for (casadi_int ri = 0; ri < k && grid; ++ri)
1124 if (idx[cj * k + ri] != (c0 + cj * cstep) * m_in + (r0 + ri * rstep))
1128 {c0, r0}, {c0 + (nc - 1) * cstep + 1, r0 + (k - 1) * rstep + 1},
1129 {0, 1}, {cstep, rstep}, node_output);
1130 work_to_onnx[o_vec[0]] = node_output;
1134 }
else if (!idx.empty()) {
1135 std::vector<casadi_int> irow = sp_in.get_row(), icol = sp_in.get_col();
1136 casadi_int nnz_in =
static_cast<casadi_int
>(irow.size());
1137 if (mx_getnonzeros.size1() == 1 && mx_getnonzeros.size2() == n_in) {
1138 casadi_int r = irow[idx[0]];
1139 std::vector<casadi_int> exp;
1140 for (casadi_int p = 0; p < nnz_in; ++p)
if (irow[p] == r) exp.push_back(p);
1144 work_to_onnx[o_vec[0]] = node_output;
1148 if (mx_getnonzeros.size1() == m_in && mx_getnonzeros.size2() == 1) {
1149 casadi_int c = icol[idx[0]];
1150 std::vector<casadi_int> exp;
1151 for (casadi_int p = 0; p < nnz_in; ++p)
if (icol[p] == c) exp.push_back(p);
1155 work_to_onnx[o_vec[0]] = node_output;
1164 emit_nonzero_remap(add_node, data, mx_getnonzeros.dep(0).sparsity(),
1165 mx_getnonzeros.sparsity(), idx, uniq, node_output);
1166 work_to_onnx[o_vec[0]] = node_output;
1178 MX mx_gnzp = f.instruction_MX(k);
1179 const Sparsity& sp_in = mx_gnzp.dep(0).sparsity();
1180 casadi_int numel_in = sp_in.size1() * sp_in.size2();
1181 std::string uniq = std::to_string(k);
1184 std::string flat = work_to_onnx[i_vec[0]];
1185 if (sp_in.size2() != 1) {
1186 flat =
"gnzp_flat_" + uniq;
1188 "gnzp_flat_" + uniq);
1192 std::string nzidx =
"gnzp_pi_" + uniq;
1193 create_cast_node(add_node, work_to_onnx[i_vec[1]], nzidx, onnx::TensorProto::INT64);
1197 std::string dpos = nzidx;
1198 if (!sp_in.is_dense()) {
1199 std::vector<casadi_int> loc = sp_in.find();
1200 std::string loc_c =
"gnzp_loc_" + uniq;
1201 add_int_constant(add_node, loc_c, loc, {1,
static_cast<casadi_int
>(loc.size())});
1202 dpos =
"gnzp_dpos_" + uniq;
1203 onnx::NodeProto* ge = add_node();
1204 ge->set_op_type(
"GatherElements");
1205 ge->add_input(loc_c); ge->add_input(nzidx); ge->add_output(dpos);
1211 onnx::NodeProto* ge = add_node();
1212 ge->set_op_type(
"GatherElements");
1213 ge->add_input(flat); ge->add_input(dpos); ge->add_output(node_output);
1216 work_to_onnx[o_vec[0]] = node_output;
1231 MX mx_snzp = f.instruction_MX(k);
1232 const Sparsity& sp_in = mx_snzp.dep(0).sparsity();
1233 casadi_int numel_in = sp_in.size1() * sp_in.size2();
1234 casadi_int np_idx = mx_snzp.dep(2).nnz();
1235 std::string uniq = std::to_string(k);
1238 std::string flat = work_to_onnx[i_vec[0]];
1239 if (sp_in.size2() != 1) {
1240 flat =
"snzp_flat_" + uniq;
1242 "snzp_flat_" + uniq);
1247 std::string nzidx =
"snzp_pi_" + uniq;
1248 create_cast_node(add_node, work_to_onnx[i_vec[2]], nzidx, onnx::TensorProto::INT64);
1249 std::string dpos = nzidx;
1250 if (!sp_in.is_dense()) {
1251 std::vector<casadi_int> loc = sp_in.find();
1252 std::string loc_c =
"snzp_loc_" + uniq;
1253 add_int_constant(add_node, loc_c, loc, {1,
static_cast<casadi_int
>(loc.size())});
1254 dpos =
"snzp_dpos_" + uniq;
1255 onnx::NodeProto* ge = add_node();
1256 ge->set_op_type(
"GatherElements");
1257 ge->add_input(loc_c); ge->add_input(nzidx); ge->add_output(dpos);
1262 std::string updates = work_to_onnx[i_vec[1]];
1263 if (mx_snzp.dep(1).size2() != 1) {
1264 updates =
"snzp_v_" + uniq;
1270 std::string scat =
"snzp_scat_" + uniq;
1271 onnx::NodeProto* se = add_node();
1272 se->set_op_type(
"ScatterElements");
1273 se->add_input(flat); se->add_input(dpos); se->add_input(updates); se->add_output(scat);
1276 onnx::AttributeProto* red = se->add_attribute();
1277 red->set_name(
"reduction");
1278 red->set_type(onnx::AttributeProto::STRING);
1283 std::string back =
"snzp_back_" + uniq;
1285 "snzp_back_" + uniq);
1286 std::string r = emit_sparsity_restore(add_node, back,
1288 mx_snzp.sparsity(),
"snzpr_" + uniq, node_output);
1289 work_to_onnx[o_vec[0]] = r;
1300 MX mx_proj = f.instruction_MX(k);
1301 Sparsity sp_out = mx_proj.sparsity();
1302 std::string r = emit_sparsity_restore(add_node, work_to_onnx[i_vec[0]],
1303 mx_proj.dep(0).sparsity(), sp_out,
1304 std::to_string(k), node_output);
1305 work_to_onnx[o_vec[0]] = r;
1320 MX mx_set = f.instruction_MX(k);
1321 Dict info = mx_set.info();
1322 std::string uniq = std::to_string(k);
1323 bool add = info[
"add"];
1325 std::vector<casadi_int> idx;
1326 if (info.count(
"nz")) {
1328 }
else if (info.count(
"slice")) {
1329 Dict s = info[
"slice"];
1330 casadi_int start = s[
"start"], stop = s[
"stop"], step = s[
"step"];
1331 for (casadi_int i = start; i < stop; i += step) idx.push_back(i);
1332 }
else if (info.count(
"inner")) {
1333 Dict ii = info[
"inner"], oo = info[
"outer"];
1334 casadi_int is = ii[
"start"], ip = ii[
"stop"], ist = ii[
"step"];
1335 casadi_int os = oo[
"start"], op2 = oo[
"stop"], ost = oo[
"step"];
1336 for (casadi_int o = os; o < op2; o += ost) {
1337 for (casadi_int i = is; i < ip; i += ist) idx.push_back(o + i);
1340 casadi_error(
"ONNX export: unsupported setnonzeros pattern.");
1342 std::vector<casadi_int> loc = mx_set.dep(0).sparsity().find();
1343 for (casadi_int& v : idx) v = loc[v];
1347 casadi_int nrow = mx_set.dep(0).size1();
1348 std::vector<casadi_int> coords;
1349 coords.reserve(2 * idx.size());
1350 for (casadi_int p : idx) { coords.push_back(p / nrow); coords.push_back(p % nrow); }
1351 std::string ind =
"snz_idx_" + uniq;
1352 add_int_constant(add_node, ind, coords, {
static_cast<casadi_int
>(idx.size()), 2});
1357 casadi_int numel_z = mx_set.dep(1).numel();
1358 std::string vals = work_to_onnx[i_vec[1]];
1359 if (mx_set.dep(1).size2() != 1) {
1360 vals =
"snz_zr_" + uniq;
1365 if (mx_set.dep(1).nnz() != numel_z) {
1366 std::vector<casadi_int> zfind = mx_set.dep(1).sparsity().find();
1367 std::string zi =
"snz_zi_" + uniq, vg =
"snz_zg_" + uniq;
1372 std::string upd =
"snz_upd_" + uniq, ush =
"snz_ush_" + uniq;
1375 node->set_op_type(
"Reshape");
1376 node->add_input(vals); node->add_input(ush); node->add_output(upd);
1378 std::string scat =
"snz_out_" + uniq;
1380 node->set_op_type(
"ScatterND");
1381 node->add_input(work_to_onnx[i_vec[0]]);
1382 node->add_input(ind);
1383 node->add_input(upd);
1384 node->add_output(scat);
1386 onnx::AttributeProto* red = node->add_attribute();
1387 red->set_name(
"reduction");
1388 red->set_type(onnx::AttributeProto::STRING);
1393 std::string r = emit_sparsity_restore(add_node, scat,
1395 mx_set.sparsity(), uniq, node_output);
1396 work_to_onnx[o_vec[0]] = r;
1402 std::string y = work_to_onnx[i_vec[0]], x = work_to_onnx[i_vec[1]];
1403 std::string p =
"atan2_" + std::to_string(k) +
"_";
1405 add_real_constant(add_node, p+
"c0", {0.0});
1406 add_real_constant(add_node, p+
"pi", {M_PI});
1420 work_to_onnx[o_vec[0]] = node_output;
1433 bool Onnx::process_operation(
1434 onnx::GraphProto* graph,
1438 const std::vector<casadi_int>& i_vec,
1439 const std::vector<casadi_int>& o_vec,
1440 std::map<casadi_int, std::string>& work_to_onnx,
1441 const std::string& node_output) {
1442 return process_operation([graph]() {
return graph->add_node(); },
1443 f, op, k, i_vec, o_vec, work_to_onnx, node_output);
static Sparsity dense(casadi_int nrow, casadi_int ncol=1)
Create a dense rectangular sparsity pattern *.
onnx::NodeProto * create_where_node(AddNodeFn add_node, const std::string &cond, const std::string &if_true, const std::string &if_false, const std::string &output)
onnx::NodeProto * create_gemm_node(AddNodeFn add_node, const std::string &A, const std::string &B, const std::string &C, const std::string &output, bool transA=false, bool transB=false)
static const OpMapping op_map[]
onnx::NodeProto * create_gather_node(AddNodeFn add_node, const std::string &data, const std::string &indices, const std::string &output, casadi_int axis=0)
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)
const OpMapping * get_op_mapping(casadi_int op)
Lookup operation mapping by CasADi opcode (for export)
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)
const OpMapping * get_op_mapping_by_name(const std::string &onnx_name)
Lookup operation mapping by ONNX name (for import)
void add_int_attribute(onnx::NodeProto *node, const std::string &name, casadi_int value)
Add an integer attribute (e.g. axis) to a node.
void emit_colmajor_reshape(AddNodeFn add_node, const std::string &data, const std::vector< casadi_int > &dims, const std::string &output, const std::string &uniq)
onnx::NodeProto * create_slice_node(AddNodeFn add_node, const std::string &uniq, const std::string &data, const std::vector< casadi_int > &starts, const std::vector< casadi_int > &ends, const std::vector< casadi_int > &axes, const std::vector< casadi_int > &steps, const std::string &output)
onnx::NodeProto * create_cast_node(AddNodeFn add_node, const std::string &input, const std::string &output, onnx::TensorProto::DataType to_type)
std::function< onnx::NodeProto *()> AddNodeFn
Callback type for adding nodes to a container (GraphProto or FunctionProto)
onnx::TensorProto * create_constant_tensor(AddNodeFn add_node, const std::string &output_name, onnx::TensorProto::DataType data_type)
void add_int_constant(onnx::GraphProto *graph, const std::string &name, const std::vector< casadi_int > &data, std::vector< casadi_int > dims={})
Add a Constant node holding an INT64 tensor to a graph (default shape: 1-D)
std::string onnx_input_name(const Function &f, casadi_int i)
Function input/output name, or a generated fallback when unnamed.
T sum(const std::vector< T > &values)
sum
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.
Operation mapping between CasADi and ONNX.