kron.cpp
1 /*
2  * This file is part of CasADi.
3  *
4  * CasADi -- A symbolic framework for dynamic optimization.
5  * Copyright (C) 2010-2023 Joel Andersson, Joris Gillis, Moritz Diehl,
6  * KU Leuven. All rights reserved.
7  * Copyright (C) 2011-2014 Greg Horn
8  *
9  * CasADi is free software; you can redistribute it and/or
10  * modify it under the terms of the GNU Lesser General Public
11  * License as published by the Free Software Foundation; either
12  * version 3 of the License, or (at your option) any later version.
13  *
14  * CasADi is distributed in the hope that it will be useful,
15  * but WITHOUT ANY WARRANTY; without even the implied warranty of
16  * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
17  * Lesser General Public License for more details.
18  *
19  * You should have received a copy of the GNU Lesser General Public
20  * License along with CasADi; if not, write to the Free Software
21  * Foundation, Inc., 51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA
22  *
23  */
24 
25 
26 #include "kron.hpp"
27 #include "casadi_misc.hpp"
28 #include "serializing_stream.hpp"
29 
30 namespace casadi {
31 
32  // ============ Sparsity-propagation traits (Fwd vs Rev) ============
33  //
34  // sp_forward and sp_reverse walk the exact same indices in the same order;
35  // they differ only in the per-iteration bit-flow direction. Captured with
36  // a `bool Fwd` template trait following the JacSparsityTraits pattern in
37  // function_internal.cpp.
38 
39  template<bool Fwd> struct KronSpTraits;
40  template<> struct KronSpTraits<true> {
41  typedef const bvec_t** arg_t;
42  typedef bvec_t** res_t;
43  static inline void step(arg_t arg, res_t res,
44  casadi_int a_el, casadi_int b_el, casadi_int k) {
45  res[0][k] = arg[0][a_el] | arg[1][b_el];
46  }
47  };
48  template<> struct KronSpTraits<false> {
49  typedef bvec_t** arg_t;
50  typedef bvec_t** res_t;
51  static inline void step(arg_t arg, res_t res,
52  casadi_int a_el, casadi_int b_el, casadi_int k) {
53  arg[0][a_el] |= res[0][k];
54  arg[1][b_el] |= res[0][k];
55  res[0][k] = 0;
56  }
57  };
58 
59  // Templated Kron sparsity walk: identical 4-loop CSC scan in both modes.
60  template<bool Fwd>
61  static int kron_sp_gen(const Sparsity& sp_a, const Sparsity& sp_b,
62  typename KronSpTraits<Fwd>::arg_t arg,
63  typename KronSpTraits<Fwd>::res_t res) {
64  const casadi_int* a_colind = sp_a.colind();
65  const casadi_int* b_colind = sp_b.colind();
66  const casadi_int a_ncol = sp_a.size2();
67  const casadi_int b_ncol = sp_b.size2();
68  casadi_int k = 0;
69  for (casadi_int a_cc=0; a_cc<a_ncol; ++a_cc) {
70  for (casadi_int b_cc=0; b_cc<b_ncol; ++b_cc) {
71  for (casadi_int a_el=a_colind[a_cc]; a_el<a_colind[a_cc+1]; ++a_el) {
72  for (casadi_int b_el=b_colind[b_cc]; b_el<b_colind[b_cc+1]; ++b_el) {
73  KronSpTraits<Fwd>::step(arg, res, a_el, b_el, k);
74  ++k;
75  }
76  }
77  }
78  }
79  return 0;
80  }
81 
82 
83  // ============ Kron (base, both-sparse) ============
84 
85  MX Kron::create(const MX& a, const MX& b) {
86  // Most-specific-first.
87  if (auto* n = DenseKron::try_create(a, b)) return MX::create(n);
88  if (auto* n = DenseSparseKron::try_create(a, b)) return MX::create(n);
89  if (auto* n = SparseDenseKron::try_create(a, b)) return MX::create(n);
90  return MX::create(new Kron(a, b));
91  }
92 
93  Kron::Kron(const MX& a, const MX& b) {
94  set_dep(a, b);
96  }
97 
98  std::string Kron::disp(const std::vector<std::string>& arg) const {
99  return "kron(" + arg.at(0) + ", " + arg.at(1) + ")";
100  }
101 
102  void Kron::eval_kernel(const double** arg, double** res) const {
103  casadi_kron(arg[0], dep(0).sparsity(), arg[1], dep(1).sparsity(), res[0]);
104  }
105 
106  void Kron::eval_kernel(const SXElem** arg, SXElem** res) const {
107  casadi_kron(arg[0], dep(0).sparsity(), arg[1], dep(1).sparsity(), res[0]);
108  }
109 
110  void Kron::eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
111  const std::vector<bool>& unique) const {
112  res[0] = kron(arg[0], arg[1]);
113  }
114 
115  int Kron::sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const {
116  return kron_sp_gen<true>(dep(0).sparsity(), dep(1).sparsity(), arg, res);
117  }
118 
119  int Kron::sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const {
120  return kron_sp_gen<false>(dep(0).sparsity(), dep(1).sparsity(), arg, res);
121  }
122 
123  void Kron::ad_forward(const std::vector<std::vector<MX> >& fseed,
124  std::vector<std::vector<MX> >& fsens) const {
125  const MX& A = dep(0);
126  const MX& B = dep(1);
127  for (casadi_int d=0; d<fsens.size(); ++d) {
128  fsens[d][0] = kron(fseed[d][0], B) + kron(A, fseed[d][1]);
129  }
130  }
131 
132  void Kron::ad_reverse(const std::vector<std::vector<MX> >& aseed,
133  std::vector<std::vector<MX> >& asens) const {
134  const MX& A = dep(0);
135  const MX& B = dep(1);
136  for (casadi_int d=0; d<aseed.size(); ++d) {
137  const MX& Fbar = aseed[d][0];
138  asens[d][0] += kron_contract(Fbar, B, true);
139  asens[d][1] += kron_contract(Fbar, A, false);
140  }
141  }
142 
144  const std::vector<casadi_int>& arg,
145  const std::vector<casadi_int>& res,
146  const std::vector<bool>& arg_is_ref,
147  std::vector<bool>& res_is_ref) const {
149  g << "casadi_kron("
150  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
151  << g.sparsity(dep(0).sparsity()) << ", "
152  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
153  << g.sparsity(dep(1).sparsity()) << ", "
154  << g.work(res[0], nnz(), false) << ");\n";
155  }
156 
159  s.pack("Kron::kind", std::string("base"));
160  }
161 
163  std::string kind;
164  s.unpack("Kron::kind", kind);
165  if (kind == "base") return new Kron(s);
166  if (kind == "dense") return new DenseKron(s);
167  if (kind == "dense_sparse") return new DenseSparseKron(s);
168  if (kind == "sparse_dense") return new SparseDenseKron(s);
169  casadi_error("Unknown Kron kind: " + kind);
170  }
171 
172 
173  // ============ DenseKron ============
174 
175  MXNode* DenseKron::try_create(const MX& a, const MX& b) {
176  if (!(a.is_dense() && b.is_dense())) return nullptr;
177  return new DenseKron(a, b);
178  }
179 
180  void DenseKron::eval_kernel(const double** arg, double** res) const {
181  casadi_kron_dense(arg[0], dep(0).size1(), dep(0).size2(),
182  arg[1], dep(1).size1(), dep(1).size2(), res[0]);
183  }
184 
185  void DenseKron::eval_kernel(const SXElem** arg, SXElem** res) const {
186  casadi_kron_dense(arg[0], dep(0).size1(), dep(0).size2(),
187  arg[1], dep(1).size1(), dep(1).size2(), res[0]);
188  }
189 
191  const std::vector<casadi_int>& arg,
192  const std::vector<casadi_int>& res,
193  const std::vector<bool>& arg_is_ref,
194  std::vector<bool>& res_is_ref) const {
196  g << "casadi_kron_dense("
197  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
198  << dep(0).size1() << ", " << dep(0).size2() << ", "
199  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
200  << dep(1).size1() << ", " << dep(1).size2() << ", "
201  << g.work(res[0], nnz(), false) << ");\n";
202  }
203 
206  s.pack("Kron::kind", std::string("dense"));
207  }
208 
209 
210  // ============ DenseSparseKron ============
211 
212  MXNode* DenseSparseKron::try_create(const MX& a, const MX& b) {
213  if (!(a.is_dense() && !b.is_dense())) return nullptr;
214  return new DenseSparseKron(a, b);
215  }
216 
217  void DenseSparseKron::eval_kernel(const double** arg, double** res) const {
218  casadi_kron_dense_sparse(arg[0], dep(0).size1(), dep(0).size2(),
219  arg[1], dep(1).sparsity(), res[0]);
220  }
221 
222  void DenseSparseKron::eval_kernel(const SXElem** arg, SXElem** res) const {
223  casadi_kron_dense_sparse(arg[0], dep(0).size1(), dep(0).size2(),
224  arg[1], dep(1).sparsity(), res[0]);
225  }
226 
228  const std::vector<casadi_int>& arg,
229  const std::vector<casadi_int>& res,
230  const std::vector<bool>& arg_is_ref,
231  std::vector<bool>& res_is_ref) const {
233  g << "casadi_kron_dense_sparse("
234  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
235  << dep(0).size1() << ", " << dep(0).size2() << ", "
236  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
237  << g.sparsity(dep(1).sparsity()) << ", "
238  << g.work(res[0], nnz(), false) << ");\n";
239  }
240 
243  s.pack("Kron::kind", std::string("dense_sparse"));
244  }
245 
246 
247  // ============ SparseDenseKron ============
248 
249  MXNode* SparseDenseKron::try_create(const MX& a, const MX& b) {
250  if (!(!a.is_dense() && b.is_dense())) return nullptr;
251  return new SparseDenseKron(a, b);
252  }
253 
254  void SparseDenseKron::eval_kernel(const double** arg, double** res) const {
256  arg[1], dep(1).size1(), dep(1).size2(), res[0]);
257  }
258 
259  void SparseDenseKron::eval_kernel(const SXElem** arg, SXElem** res) const {
261  arg[1], dep(1).size1(), dep(1).size2(), res[0]);
262  }
263 
265  const std::vector<casadi_int>& arg,
266  const std::vector<casadi_int>& res,
267  const std::vector<bool>& arg_is_ref,
268  std::vector<bool>& res_is_ref) const {
270  g << "casadi_kron_sparse_dense("
271  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
272  << g.sparsity(dep(0).sparsity()) << ", "
273  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
274  << dep(1).size1() << ", " << dep(1).size2() << ", "
275  << g.work(res[0], nnz(), false) << ");\n";
276  }
277 
280  s.pack("Kron::kind", std::string("sparse_dense"));
281  }
282 
283 
284  // ============ KronContract (base, fully general) ============
285 
286  MX KronContract::create(const MX& m, const MX& x, bool inner) {
287  if (auto* n = DenseKronContract::try_create(m, x, inner)) return MX::create(n);
288  if (auto* n = DenseSparseKronContract::try_create(m, x, inner)) return MX::create(n);
289  if (auto* n = SparseDenseKronContract::try_create(m, x, inner)) return MX::create(n);
290  return MX::create(new KronContract(m, x, inner));
291  }
292 
293  KronContract::KronContract(const MX& m, const MX& x, bool inner) : inner_(inner) {
294  casadi_assert(x.size1() > 0 && x.size2() > 0,
295  "KronContract: X must be nonempty");
296  casadi_assert(m.size1() % x.size1() == 0 && m.size2() % x.size2() == 0,
297  "KronContract: M dims must be multiples of X dims");
298  set_dep(m, x);
300  }
301 
302  std::string KronContract::disp(const std::vector<std::string>& arg) const {
303  return std::string("kron_contract(") + arg.at(0) + ", " + arg.at(1)
304  + ", " + (inner_ ? "inner" : "outer") + ")";
305  }
306 
307  size_t KronContract::sz_w() const {
308  return static_cast<size_t>(dep(1).numel());
309  }
310 
311  void KronContract::eval_kernel(const double** arg, double** res, double* w) const {
312  if (inner_) {
314  arg[1], dep(1).sparsity(),
315  res[0], sparsity(), w);
316  } else {
318  arg[1], dep(1).sparsity(),
319  res[0], sparsity(), w);
320  }
321  }
322 
323  void KronContract::eval_kernel(const SXElem** arg, SXElem** res, SXElem* w) const {
324  if (inner_) {
326  arg[1], dep(1).sparsity(),
327  res[0], sparsity(), w);
328  } else {
330  arg[1], dep(1).sparsity(),
331  res[0], sparsity(), w);
332  }
333  }
334 
335  void KronContract::eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
336  const std::vector<bool>& unique) const {
337  res[0] = kron_contract(arg[0], arg[1], inner_);
338  }
339 
340  int KronContract::sp_forward(const bvec_t** arg, bvec_t** res,
341  casadi_int* iw, bvec_t* w) const {
342  // sp_forward and sp_reverse here have asymmetric X access (forward reads
343  // X as input; reverse writes X-bar as output via per-iteration X-CSC
344  // scan). Unifying via a shared template would require staging X-bar
345  // through w in reverse and scattering at the end -- a real algorithmic
346  // shift that can slow the (very sparse M, very dense X) corner by 100x+.
347  // We keep them as two functions; the common scaffold is the column-
348  // decomposition + row decomposition pattern which is short enough to
349  // duplicate readably.
350  const Sparsity& sp_m = dep(0).sparsity();
351  const Sparsity& sp_x = dep(1).sparsity();
352  const Sparsity& sp_y = sparsity();
353  const casadi_int xrow = sp_x.size1(), xcol = sp_x.size2();
354  const casadi_int yrow = sp_y.size1(), ycol = sp_y.size2();
355  const casadi_int* m_colind = sp_m.colind();
356  const casadi_int* m_row = sp_m.row();
357  const casadi_int* x_colind = sp_x.colind();
358  const casadi_int* x_row = sp_x.row();
359  const casadi_int* y_colind = sp_y.colind();
360  const casadi_int* y_row = sp_y.row();
361  const casadi_int mB = inner_ ? xrow : yrow;
362  const casadi_int nB = inner_ ? xcol : ycol;
363  const casadi_int nA = inner_ ? ycol : xcol;
364  bvec_t* w_x = w;
365  const casadi_int xn = xrow*xcol;
366  for (casadi_int k=0; k<xn; ++k) w_x[k] = 0;
367  for (casadi_int cc=0; cc<xcol; ++cc) {
368  for (casadi_int el=x_colind[cc]; el<x_colind[cc+1]; ++el) {
369  w_x[cc*xrow + x_row[el]] = arg[1][el];
370  }
371  }
372  const casadi_int yn_total = y_colind[ycol];
373  for (casadi_int k=0; k<yn_total; ++k) res[0][k] = 0;
374  for (casadi_int j=0; j<nA; ++j) {
375  for (casadi_int s=0; s<nB; ++s) {
376  casadi_int cc = j*nB + s;
377  casadi_int y_cc = inner_ ? j : s;
378  casadi_int x_cc = inner_ ? s : j;
379  casadi_int y_col_start = y_colind[y_cc];
380  casadi_int y_col_end = y_colind[y_cc+1];
381  if (y_col_start == y_col_end) continue;
382  for (casadi_int el=m_colind[cc]; el<m_colind[cc+1]; ++el) {
383  casadi_int rr = m_row[el];
384  casadi_int outer_row = rr / mB;
385  casadi_int inner_row = rr % mB;
386  casadi_int y_rr = inner_ ? outer_row : inner_row;
387  casadi_int x_rr = inner_ ? inner_row : outer_row;
388  bvec_t x_pat = w_x[x_cc*xrow + x_rr];
389  for (casadi_int y_el=y_col_start; y_el<y_col_end; ++y_el) {
390  if (y_row[y_el] == y_rr) { res[0][y_el] |= arg[0][el] | x_pat; break; }
391  if (y_row[y_el] > y_rr) break;
392  }
393  }
394  }
395  }
396  return 0;
397  }
398 
400  casadi_int* iw, bvec_t* w) const {
401  const Sparsity& sp_m = dep(0).sparsity();
402  const Sparsity& sp_x = dep(1).sparsity();
403  const Sparsity& sp_y = sparsity();
404  const casadi_int xrow = sp_x.size1(), xcol = sp_x.size2();
405  const casadi_int yrow = sp_y.size1(), ycol = sp_y.size2();
406  const casadi_int* m_colind = sp_m.colind();
407  const casadi_int* m_row = sp_m.row();
408  const casadi_int* x_colind = sp_x.colind();
409  const casadi_int* x_row = sp_x.row();
410  const casadi_int* y_colind = sp_y.colind();
411  const casadi_int* y_row = sp_y.row();
412  const casadi_int mB = inner_ ? xrow : yrow;
413  const casadi_int nB = inner_ ? xcol : ycol;
414  const casadi_int nA = inner_ ? ycol : xcol;
415  for (casadi_int j=0; j<nA; ++j) {
416  for (casadi_int s=0; s<nB; ++s) {
417  casadi_int cc = j*nB + s;
418  casadi_int y_cc = inner_ ? j : s;
419  casadi_int x_cc = inner_ ? s : j;
420  casadi_int y_col_start = y_colind[y_cc];
421  casadi_int y_col_end = y_colind[y_cc+1];
422  if (y_col_start == y_col_end) continue;
423  casadi_int x_col_start = x_colind[x_cc];
424  casadi_int x_col_end = x_colind[x_cc+1];
425  for (casadi_int el=m_colind[cc]; el<m_colind[cc+1]; ++el) {
426  casadi_int rr = m_row[el];
427  casadi_int outer_row = rr / mB;
428  casadi_int inner_row = rr % mB;
429  casadi_int y_rr = inner_ ? outer_row : inner_row;
430  casadi_int x_rr = inner_ ? inner_row : outer_row;
431  for (casadi_int y_el=y_col_start; y_el<y_col_end; ++y_el) {
432  if (y_row[y_el] == y_rr) {
433  bvec_t sval = res[0][y_el];
434  arg[0][el] |= sval;
435  for (casadi_int x_el=x_col_start; x_el<x_col_end; ++x_el) {
436  if (x_row[x_el] == x_rr) { arg[1][x_el] |= sval; break; }
437  if (x_row[x_el] > x_rr) break;
438  }
439  break;
440  }
441  if (y_row[y_el] > y_rr) break;
442  }
443  }
444  }
445  }
446  const casadi_int yn_total = y_colind[ycol];
447  for (casadi_int k=0; k<yn_total; ++k) res[0][k] = 0;
448  return 0;
449  }
450 
451  void KronContract::ad_forward(const std::vector<std::vector<MX> >& fseed,
452  std::vector<std::vector<MX> >& fsens) const {
453  const MX& M = dep(0);
454  const MX& X = dep(1);
455  for (casadi_int d=0; d<fsens.size(); ++d) {
456  fsens[d][0] = kron_contract(fseed[d][0], X, inner_)
457  + kron_contract(M, fseed[d][1], inner_);
458  }
459  }
460 
461  void KronContract::ad_reverse(const std::vector<std::vector<MX> >& aseed,
462  std::vector<std::vector<MX> >& asens) const {
463  const MX& M = dep(0);
464  const MX& X = dep(1);
465  for (casadi_int d=0; d<aseed.size(); ++d) {
466  const MX& Ybar = aseed[d][0];
467  MX m_kron = inner_ ? kron(Ybar, X) : kron(X, Ybar);
468  asens[d][0] += project(m_kron, M.sparsity());
469  asens[d][1] += kron_contract(M, Ybar, !inner_);
470  }
471  }
472 
474  const std::vector<casadi_int>& arg,
475  const std::vector<casadi_int>& res,
476  const std::vector<bool>& arg_is_ref,
477  std::vector<bool>& res_is_ref) const {
478  if (inner_) {
480  g << "casadi_kron_contract_inner(";
481  } else {
483  g << "casadi_kron_contract_outer(";
484  }
485  g << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
486  << g.sparsity(dep(0).sparsity()) << ", "
487  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
488  << g.sparsity(dep(1).sparsity()) << ", "
489  << g.work(res[0], nnz(), false) << ", "
490  << g.sparsity(sparsity()) << ", w);\n";
491  }
492 
495  s.pack("KronContract::inner", inner_);
496  }
497 
500  s.pack("KronContract::kind", std::string("base"));
501  }
502 
504  s.unpack("KronContract::inner", inner_);
505  }
506 
508  std::string kind;
509  s.unpack("KronContract::kind", kind);
510  if (kind == "base") return new KronContract(s);
511  if (kind == "dense") return new DenseKronContract(s);
512  if (kind == "dense_sparse") return new DenseSparseKronContract(s);
513  if (kind == "sparse_dense") return new SparseDenseKronContract(s);
514  casadi_error("Unknown KronContract kind: " + kind);
515  }
516 
517 
518  // ============ DenseKronContract ============
519 
520  MXNode* DenseKronContract::try_create(const MX& m, const MX& x, bool inner) {
521  if (!(m.is_dense() && x.is_dense())) return nullptr;
522  return new DenseKronContract(m, x, inner);
523  }
524 
525  void DenseKronContract::eval_kernel(const double** arg, double** res, double* w) const {
526  if (inner_) {
527  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
528  const casadi_int mB = dep(1).size1(), nB = dep(1).size2();
529  casadi_kron_contract_inner_dense(arg[0], mA, nA, arg[1], mB, nB, res[0]);
530  } else {
531  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
532  const casadi_int mA = dep(1).size1(), nA = dep(1).size2();
533  casadi_kron_contract_outer_dense(arg[0], mB, nB, arg[1], mA, nA, res[0]);
534  }
535  }
536 
537  void DenseKronContract::eval_kernel(const SXElem** arg, SXElem** res, SXElem* w) const {
538  if (inner_) {
539  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
540  const casadi_int mB = dep(1).size1(), nB = dep(1).size2();
541  casadi_kron_contract_inner_dense(arg[0], mA, nA, arg[1], mB, nB, res[0]);
542  } else {
543  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
544  const casadi_int mA = dep(1).size1(), nA = dep(1).size2();
545  casadi_kron_contract_outer_dense(arg[0], mB, nB, arg[1], mA, nA, res[0]);
546  }
547  }
548 
550  const std::vector<casadi_int>& arg,
551  const std::vector<casadi_int>& res,
552  const std::vector<bool>& arg_is_ref,
553  std::vector<bool>& res_is_ref) const {
554  if (inner_) {
556  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
557  const casadi_int mB = dep(1).size1(), nB = dep(1).size2();
558  g << "casadi_kron_contract_inner_dense("
559  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
560  << mA << ", " << nA << ", "
561  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
562  << mB << ", " << nB << ", "
563  << g.work(res[0], nnz(), false) << ");\n";
564  } else {
566  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
567  const casadi_int mA = dep(1).size1(), nA = dep(1).size2();
568  g << "casadi_kron_contract_outer_dense("
569  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
570  << mB << ", " << nB << ", "
571  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
572  << mA << ", " << nA << ", "
573  << g.work(res[0], nnz(), false) << ");\n";
574  }
575  }
576 
579  s.pack("KronContract::kind", std::string("dense"));
580  }
581 
582 
583  // ============ DenseSparseKronContract ============
584 
585  MXNode* DenseSparseKronContract::try_create(const MX& m, const MX& x, bool inner) {
586  if (!(m.is_dense() && !x.is_dense())) return nullptr;
587  return new DenseSparseKronContract(m, x, inner);
588  }
589 
590  void DenseSparseKronContract::eval_kernel(const double** arg, double** res, double* w) const {
591  if (inner_) {
592  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
594  arg[1], dep(1).sparsity(), res[0]);
595  } else {
596  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
598  arg[1], dep(1).sparsity(), res[0]);
599  }
600  }
601 
602  void DenseSparseKronContract::eval_kernel(const SXElem** arg, SXElem** res, SXElem* w) const {
603  if (inner_) {
604  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
606  arg[1], dep(1).sparsity(), res[0]);
607  } else {
608  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
610  arg[1], dep(1).sparsity(), res[0]);
611  }
612  }
613 
615  const std::vector<casadi_int>& arg,
616  const std::vector<casadi_int>& res,
617  const std::vector<bool>& arg_is_ref,
618  std::vector<bool>& res_is_ref) const {
619  if (inner_) {
621  const casadi_int mA = sparsity().size1(), nA = sparsity().size2();
622  g << "casadi_kron_contract_inner_dense_sparse("
623  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
624  << mA << ", " << nA << ", "
625  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
626  << g.sparsity(dep(1).sparsity()) << ", "
627  << g.work(res[0], nnz(), false) << ");\n";
628  } else {
630  const casadi_int mB = sparsity().size1(), nB = sparsity().size2();
631  g << "casadi_kron_contract_outer_dense_sparse("
632  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
633  << mB << ", " << nB << ", "
634  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
635  << g.sparsity(dep(1).sparsity()) << ", "
636  << g.work(res[0], nnz(), false) << ");\n";
637  }
638  }
639 
642  s.pack("KronContract::kind", std::string("dense_sparse"));
643  }
644 
645 
646  // ============ SparseDenseKronContract ============
647 
648  MXNode* SparseDenseKronContract::try_create(const MX& m, const MX& x, bool inner) {
649  if (!(!m.is_dense() && x.is_dense())) return nullptr;
650  return new SparseDenseKronContract(m, x, inner);
651  }
652 
653  void SparseDenseKronContract::eval_kernel(const double** arg, double** res, double* w) const {
654  if (inner_) {
656  arg[1], dep(1).size1(), dep(1).size2(),
657  res[0], sparsity());
658  } else {
660  arg[1], dep(1).size1(), dep(1).size2(),
661  res[0], sparsity());
662  }
663  }
664 
665  void SparseDenseKronContract::eval_kernel(const SXElem** arg, SXElem** res, SXElem* w) const {
666  if (inner_) {
668  arg[1], dep(1).size1(), dep(1).size2(),
669  res[0], sparsity());
670  } else {
672  arg[1], dep(1).size1(), dep(1).size2(),
673  res[0], sparsity());
674  }
675  }
676 
678  const std::vector<casadi_int>& arg,
679  const std::vector<casadi_int>& res,
680  const std::vector<bool>& arg_is_ref,
681  std::vector<bool>& res_is_ref) const {
682  if (inner_) {
684  g << "casadi_kron_contract_inner_sparse_dense("
685  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
686  << g.sparsity(dep(0).sparsity()) << ", "
687  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
688  << dep(1).size1() << ", " << dep(1).size2() << ", "
689  << g.work(res[0], nnz(), false) << ", "
690  << g.sparsity(sparsity()) << ");\n";
691  } else {
693  g << "casadi_kron_contract_outer_sparse_dense("
694  << g.work(arg[0], dep(0).nnz(), arg_is_ref[0]) << ", "
695  << g.sparsity(dep(0).sparsity()) << ", "
696  << g.work(arg[1], dep(1).nnz(), arg_is_ref[1]) << ", "
697  << dep(1).size1() << ", " << dep(1).size2() << ", "
698  << g.work(res[0], nnz(), false) << ", "
699  << g.sparsity(sparsity()) << ");\n";
700  }
701  }
702 
705  s.pack("KronContract::kind", std::string("sparse_dense"));
706  }
707 
708 } // namespace casadi
Helper class for C code generation.
std::string work(casadi_int n, casadi_int sz, bool is_ref) const
std::string sparsity(const Sparsity &sp, bool canonical=true)
void add_auxiliary(Auxiliary f, const std::vector< std::string > &inst={"casadi_real"})
Add a built-in auxiliary function.
KronContract specialization: M dense, X dense (=> Y dense)
Definition: kron.hpp:361
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:577
DenseKronContract(const MX &m, const MX &x, bool inner)
Definition: kron.hpp:365
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:549
static MXNode * try_create(const MX &m, const MX &x, bool inner)
Definition: kron.cpp:520
void eval_kernel(const double **arg, double **res, double *w) const override
Definition: kron.cpp:525
Kron specialization: both operands dense.
Definition: kron.hpp:155
void eval_kernel(const double **arg, double **res) const override
Subclass hook: r = kron(a, b) on the input buffers.
Definition: kron.cpp:180
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:204
static MXNode * try_create(const MX &a, const MX &b)
Definition: kron.cpp:175
DenseKron(const MX &a, const MX &b)
Definition: kron.hpp:159
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:190
KronContract specialization: M dense, X sparse (=> Y dense)
Definition: kron.hpp:388
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:614
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:640
DenseSparseKronContract(const MX &m, const MX &x, bool inner)
Definition: kron.hpp:392
void eval_kernel(const double **arg, double **res, double *w) const override
Definition: kron.cpp:590
static MXNode * try_create(const MX &m, const MX &x, bool inner)
Definition: kron.cpp:585
Kron specialization: dense a + sparse b.
Definition: kron.hpp:179
DenseSparseKron(const MX &a, const MX &b)
Definition: kron.hpp:183
void eval_kernel(const double **arg, double **res) const override
Subclass hook: r = kron(a, b) on the input buffers.
Definition: kron.cpp:217
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:227
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:241
static MXNode * try_create(const MX &a, const MX &b)
Definition: kron.cpp:212
Helper class for Serialization.
void unpack(Sparsity &e)
Reconstruct an object from the input stream.
casadi_int numel() const
Get the number of elements.
bool is_dense() const
Check if the matrix expression is dense.
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)
void ad_reverse(const std::vector< std::vector< MX > > &aseed, std::vector< std::vector< MX > > &asens) const override
Calculate reverse mode directional derivatives.
Definition: kron.cpp:461
int sp_reverse(bvec_t **arg, bvec_t **res, casadi_int *iw, bvec_t *w) const override
Propagate sparsity backwards.
Definition: kron.cpp:399
void eval_mx(const std::vector< MX > &arg, std::vector< MX > &res, const std::vector< bool > &unique={}) const override
Evaluate symbolically (MX)
Definition: kron.cpp:335
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:473
void serialize_body(SerializingStream &s) const override
Serialize an object without type information.
Definition: kron.cpp:493
static MX create(const MX &m, const MX &x, bool inner)
Factory: dispatch to the most specific subclass for the given operands.
Definition: kron.cpp:286
int sp_forward(const bvec_t **arg, bvec_t **res, casadi_int *iw, bvec_t *w) const override
Propagate sparsity forward.
Definition: kron.cpp:340
void ad_forward(const std::vector< std::vector< MX > > &fseed, std::vector< std::vector< MX > > &fsens) const override
Calculate forward mode directional derivatives.
Definition: kron.cpp:451
std::string disp(const std::vector< std::string > &arg) const override
Print expression.
Definition: kron.cpp:302
virtual void eval_kernel(const double **arg, double **res, double *w) const
Definition: kron.cpp:311
static MXNode * deserialize(DeserializingStream &s)
Deserialize with type disambiguation.
Definition: kron.cpp:507
bool inner_
Which axes are contracted (true = inner (mB, nB); false = outer (mA, nA))
Definition: kron.hpp:347
KronContract(const MX &m, const MX &x, bool inner)
Constructor (auto-derives output sparsity)
Definition: kron.cpp:293
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:498
size_t sz_w() const override
Get required length of w field.
Definition: kron.cpp:307
static MXNode * deserialize(DeserializingStream &s)
Deserialize with type disambiguation.
Definition: kron.cpp:162
Kron(const MX &a, const MX &b)
Constructor.
Definition: kron.cpp:93
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:157
std::string disp(const std::vector< std::string > &arg) const override
Print expression.
Definition: kron.cpp:98
static MX create(const MX &a, const MX &b)
Factory: dispatch to the most specific subclass for the given operands.
Definition: kron.cpp:85
int sp_reverse(bvec_t **arg, bvec_t **res, casadi_int *iw, bvec_t *w) const override
Propagate sparsity backwards.
Definition: kron.cpp:119
virtual void eval_kernel(const double **arg, double **res) const
Subclass hook: r = kron(a, b) on the input buffers.
Definition: kron.cpp:102
int sp_forward(const bvec_t **arg, bvec_t **res, casadi_int *iw, bvec_t *w) const override
Propagate sparsity forward.
Definition: kron.cpp:115
void eval_mx(const std::vector< MX > &arg, std::vector< MX > &res, const std::vector< bool > &unique={}) const override
Evaluate symbolically (MX)
Definition: kron.cpp:110
void ad_reverse(const std::vector< std::vector< MX > > &aseed, std::vector< std::vector< MX > > &asens) const override
Calculate reverse mode directional derivatives.
Definition: kron.cpp:132
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:143
void ad_forward(const std::vector< std::vector< MX > > &fseed, std::vector< std::vector< MX > > &fsens) const override
Calculate forward mode directional derivatives.
Definition: kron.cpp:123
Node class for MX objects.
Definition: mx_node.hpp:51
virtual void serialize_type(SerializingStream &s) const
Serialize type information.
Definition: mx_node.cpp:535
const Sparsity & sparsity() const
Get the sparsity.
Definition: mx_node.hpp:410
casadi_int size2() const
Definition: mx_node.hpp:429
casadi_int nnz(casadi_int i=0) const
Definition: mx_node.hpp:427
const MX & dep(casadi_int ind=0) const
dependencies - functions that have to be evaluated before this one
Definition: mx_node.hpp:392
virtual void serialize_body(SerializingStream &s) const
Serialize an object without type information.
Definition: mx_node.cpp:530
void set_sparsity(const Sparsity &sparsity)
Set the sparsity.
Definition: mx_node.cpp:224
casadi_int size1() const
Definition: mx_node.hpp:428
void set_dep(const MX &dep)
Set unary dependency.
Definition: mx_node.cpp:228
MX - Matrix expression.
Definition: mx.hpp:92
static MX create(MXNode *node)
Create from node.
Definition: mx.cpp:69
const Sparsity & sparsity() const
Get the sparsity pattern.
Definition: mx.cpp:612
The basic scalar symbolic class of CasADi.
Definition: sx_elem.hpp:75
Helper class for Serialization.
void pack(const Sparsity &e)
Serializes an object to the output stream.
KronContract specialization: M sparse, X dense.
Definition: kron.hpp:414
SparseDenseKronContract(const MX &m, const MX &x, bool inner)
Definition: kron.hpp:418
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:703
void eval_kernel(const double **arg, double **res, double *w) const override
Definition: kron.cpp:653
static MXNode * try_create(const MX &m, const MX &x, bool inner)
Definition: kron.cpp:648
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:677
Kron specialization: sparse a + dense b.
Definition: kron.hpp:203
void serialize_type(SerializingStream &s) const override
Serialize specific part of node.
Definition: kron.cpp:278
SparseDenseKron(const MX &a, const MX &b)
Definition: kron.hpp:207
static MXNode * try_create(const MX &a, const MX &b)
Definition: kron.cpp:249
void generate(CodeGenerator &g, const std::vector< casadi_int > &arg, const std::vector< casadi_int > &res, const std::vector< bool > &arg_is_ref, std::vector< bool > &res_is_ref) const override
Generate code for the operation.
Definition: kron.cpp:264
void eval_kernel(const double **arg, double **res) const override
Subclass hook: r = kron(a, b) on the input buffers.
Definition: kron.cpp:254
General sparsity class.
Definition: sparsity.hpp:106
casadi_int size1() const
Get the number of rows.
Definition: sparsity.cpp:124
casadi_int size2() const
Get the number of columns.
Definition: sparsity.cpp:128
const casadi_int * row() const
Get a reference to row-vector,.
Definition: sparsity.cpp:164
static Sparsity kron(const Sparsity &a, const Sparsity &b)
Enlarge matrix.
Definition: sparsity.cpp:1450
const casadi_int * colind() const
Get a reference to the colindex of all column element (see class description)
Definition: sparsity.cpp:168
static Sparsity kron_contract(const Sparsity &sp_m, const Sparsity &sp_x, bool inner)
Output sparsity of casadi::KronContract.
Definition: sparsity.cpp:1494
The casadi namespace.
Definition: archiver.cpp:28
void casadi_kron_contract_inner(const T1 *m, const casadi_int *sp_m, const T1 *b, const casadi_int *sp_b, T1 *y, const casadi_int *sp_y, T1 *w)
static int kron_sp_gen(const Sparsity &sp_a, const Sparsity &sp_b, typename KronSpTraits< Fwd >::arg_t arg, typename KronSpTraits< Fwd >::res_t res)
Definition: kron.cpp:61
void casadi_kron_contract_inner_dense(const T1 *m, casadi_int mA, casadi_int nA, const T1 *b, casadi_int mB, casadi_int nB, T1 *y)
void casadi_kron_sparse_dense(const T1 *a, const casadi_int *sp_a, const T1 *b, casadi_int mB, casadi_int nB, T1 *r)
unsigned long long bvec_t
void casadi_kron_contract_outer_dense(const T1 *m, casadi_int mB, casadi_int nB, const T1 *a, casadi_int mA, casadi_int nA, T1 *y)
void casadi_kron_contract_inner_sparse_dense(const T1 *m, const casadi_int *sp_m, const T1 *b, casadi_int mB, casadi_int nB, T1 *y, const casadi_int *sp_y)
void casadi_kron_contract_inner_dense_sparse(const T1 *m, casadi_int mA, casadi_int nA, const T1 *b, const casadi_int *sp_b, T1 *y)
void casadi_kron_contract_outer_dense_sparse(const T1 *m, casadi_int mB, casadi_int nB, const T1 *a, const casadi_int *sp_a, T1 *y)
void casadi_kron(const T1 *a, const casadi_int *sp_a, const T1 *b, const casadi_int *sp_b, T1 *r)
MX kron_contract(const MX &m, const MX &x, bool inner)
Kronecker contraction.
Definition: mx.hpp:1121
void casadi_kron_contract_outer_sparse_dense(const T1 *m, const casadi_int *sp_m, const T1 *a, casadi_int mA, casadi_int nA, T1 *y, const casadi_int *sp_y)
void casadi_kron_dense(const T1 *a, casadi_int mA, casadi_int nA, const T1 *b, casadi_int mB, casadi_int nB, T1 *r)
void casadi_kron_contract_outer(const T1 *m, const casadi_int *sp_m, const T1 *a, const casadi_int *sp_a, T1 *y, const casadi_int *sp_y, T1 *w)
void casadi_kron_dense_sparse(const T1 *a, casadi_int mA, casadi_int nA, const T1 *b, const casadi_int *sp_b, T1 *r)