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;
95 casadi_int n_instr = f.n_instructions();
98 std::map<casadi_int, casadi_int> output_segment_count;
99 for (casadi_int k = 0; k < n_instr; ++k) {
101 MX mx = f.instruction_MX(k);
102 Dict info = mx.info();
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) {
114 casadi_int op = f.instruction_id(k);
115 std::vector<casadi_int> o = f.instruction_output(k);
116 std::vector<casadi_int> i = f.instruction_input(k);
119 std::string node_output =
"n" + std::to_string(k);
123 MX mx = f.instruction_MX(k);
124 Dict info = mx.info();
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();
144 Sparsity out_sp = f.sparsity_out(output_idx);
145 auto add_node = [graph]() -> onnx::NodeProto* {
return graph->add_node(); };
146 bool folded = (dep.size1() != out_sp.size1() || dep.size2() != out_sp.size2());
147 if (folded && !out_sp.is_dense()) {
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;
166 Function called_func = f.instruction_MX(k).which_function();
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];
201 casadi_int R = f.size1_out(output_idx),
C = f.size2_out(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;
213 Sparsity out_sp = f.sparsity_out(output_idx);
215 std::string uniq =
"oseg" + std::to_string(output_idx);
218 auto finish = [&](
const std::function<void(
const std::string&)>& emit) ->
void {
219 if (out_sp.is_dense()) {
222 std::string tmp = oname +
"_pre";
224 emit_sparsity_restore(seg_add_node, tmp,
Sparsity::dense(R, C), out_sp,
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;
269 for (casadi_int i = 0; i < f.n_out(); ++i) output_names.insert(
onnx_output_name(f, i));
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)) {
328 const Sparsity& sp = f.sparsity_out(i);
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) {
347 if (!f.sparsity_in(i).is_dense()) {
348 fill_sparse_tensor(graph->add_sparse_initializer(),
onnx_input_name(f, i),
349 DM(f.sparsity_in(i), 0.0));
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;
std::set< std::string > exported_functions_
Track which functions have been exported as FunctionProto.
static Sparsity dense(casadi_int nrow, casadi_int ncol=1)
Create a dense rectangular sparsity pattern *.
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
std::string onnx_output_name(const Function &f, casadi_int i)
void add_int_attribute(onnx::NodeProto *node, const std::string &name, casadi_int value)
Add an integer attribute (e.g. axis) to a node.
std::string onnx_input_name(const Function &f, casadi_int i)
Function input/output name, or a generated fallback when unnamed.