getnonzeros.hpp
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 #ifndef CASADI_GETNONZEROS_HPP
27 #define CASADI_GETNONZEROS_HPP
28 
29 #include "mx_node.hpp"
30 #include <map>
31 #include <stack>
32 
34 
35 namespace casadi {
42  class CASADI_EXPORT GetNonzeros : public MXNode {
43  public:
44 
47  static MX create(const Sparsity& sp, const MX& x, const std::vector<casadi_int>& nz);
48  static MX create(const Sparsity& sp, const MX& x, const Slice& s);
49  static MX create(const Sparsity& sp, const MX& x, const Slice& inner, const Slice& outer);
51 
53  GetNonzeros(const Sparsity& sp, const MX& y);
54 
56  ~GetNonzeros() override {}
57 
61  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
62  const std::vector<bool>& unique={}) const override;
63 
67  void eval_linear(const std::vector<std::array<MX, 3> >& arg,
68  std::vector<std::array<MX, 3> >& res) const override;
69 
73  int eval_activity(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override {
74  return sp_forward(arg, res, iw, w);
75  }
76 
80  void ad_forward(const std::vector<std::vector<MX> >& fseed,
81  std::vector<std::vector<MX> >& fsens) const override;
82 
86  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
87  std::vector<std::vector<MX> >& asens) const override;
88 
90  Matrix<casadi_int> mapping() const override;
91 
93  virtual std::vector<casadi_int> all() const = 0;
94 
98  casadi_int op() const override { return OP_GETNONZEROS;}
99 
101  MX get_nzref(const Sparsity& sp, const std::vector<casadi_int>& nz,
102  bool unique=false) const override;
103 
107  static MXNode* deserialize(DeserializingStream& s);
108 
109  protected:
113  explicit GetNonzeros(DeserializingStream& s) : MXNode(s) {}
114  };
115 
116  class CASADI_EXPORT GetNonzerosVector : public GetNonzeros {
117  public:
119  GetNonzerosVector(const Sparsity& sp, const MX& x,
120  const std::vector<casadi_int>& nz) : GetNonzeros(sp, x), nz_(nz) {}
121 
123  ~GetNonzerosVector() override {}
124 
126  std::vector<casadi_int> all() const override { return nz_;}
127 
131  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
132 
136  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
137 
141  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
142  const std::vector<bool>& unique={}) const override;
143 
145  template<typename T>
146  int eval_gen(const T* const* arg, T* const* res, casadi_int* iw, T* w) const;
147 
149  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
150 
152  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
153 
157  std::string disp(const std::vector<std::string>& arg) const override;
158 
162  void generate(CodeGenerator& g,
163  const std::vector<casadi_int>& arg,
164  const std::vector<casadi_int>& res,
165  const std::vector<bool>& arg_is_ref,
166  std::vector<bool>& res_is_ref) const override;
167 
171  bool is_equal(const MXNode* node, casadi_int depth) const override;
172 
174  Dict info() const override { return {{"nz", nz_}}; }
175 
177  std::vector<casadi_int> nz_;
178 
182  void serialize_body(SerializingStream& s) const override;
186  void serialize_type(SerializingStream& s) const override;
187 
191  explicit GetNonzerosVector(DeserializingStream& s);
192  };
193 
194  // Specialization of the above when nz_ is a Slice
195  class CASADI_EXPORT GetNonzerosSlice : public GetNonzeros {
196  public:
197 
199  GetNonzerosSlice(const Sparsity& sp, const MX& x, const Slice& s) : GetNonzeros(sp, x), s_(s) {}
200 
202  ~GetNonzerosSlice() override {}
203 
205  std::vector<casadi_int> all() const override { return s_.all(s_.stop);}
206 
210  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
211 
215  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
216 
218  template<typename T>
219  int eval_gen(const T* const* arg, T* const* res, casadi_int* iw, T* w) const;
220 
222  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
223 
225  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
226 
230  std::string disp(const std::vector<std::string>& arg) const override;
231 
235  void generate(CodeGenerator& g,
236  const std::vector<casadi_int>& arg,
237  const std::vector<casadi_int>& res,
238  const std::vector<bool>& arg_is_ref,
239  std::vector<bool>& res_is_ref) const override;
240 
244  bool is_equal(const MXNode* node, casadi_int depth) const override;
245 
247  Dict info() const override { return {{"slice", s_.info()}}; }
248 
249  // Data member
250  Slice s_;
251 
255  void serialize_body(SerializingStream& s) const override;
259  void serialize_type(SerializingStream& s) const override;
260 
264  explicit GetNonzerosSlice(DeserializingStream& s);
265  };
266 
267  // Specialization of the above when nz_ is a nested Slice
268  class CASADI_EXPORT GetNonzerosSlice2 : public GetNonzeros {
269  public:
270 
272  GetNonzerosSlice2(const Sparsity& sp, const MX& x, const Slice& inner,
273  const Slice& outer) : GetNonzeros(sp, x), inner_(inner), outer_(outer) {}
274 
276  ~GetNonzerosSlice2() override {}
277 
279  std::vector<casadi_int> all() const override { return inner_.all(outer_, outer_.stop);}
280 
284  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
285 
289  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
290 
292  template<typename T>
293  int eval_gen(const T* const* arg, T* const* res, casadi_int* iw, T* w) const;
294 
296  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
297 
299  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
300 
304  std::string disp(const std::vector<std::string>& arg) const override;
305 
309  void generate(CodeGenerator& g,
310  const std::vector<casadi_int>& arg,
311  const std::vector<casadi_int>& res,
312  const std::vector<bool>& arg_is_ref,
313  std::vector<bool>& res_is_ref) const override;
314 
318  bool is_equal(const MXNode* node, casadi_int depth) const override;
319 
321  Dict info() const override { return {{"inner", inner_.info()}, {"outer", outer_.info()}}; }
322 
323  // Data members
324  Slice inner_, outer_;
325 
329  void serialize_body(SerializingStream& s) const override;
333  void serialize_type(SerializingStream& s) const override;
334 
338  explicit GetNonzerosSlice2(DeserializingStream& s);
339  };
340 
341 
342 } // namespace casadi
344 
345 #endif // CASADI_GETNONZEROS_HPP
The casadi namespace.
Definition: archiver.hpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.