setnonzeros.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_SETNONZEROS_HPP
27 #define CASADI_SETNONZEROS_HPP
28 
29 #include "mx_node.hpp"
30 #include <map>
31 #include <stack>
32 
34 
35 namespace casadi {
36 
43  template<bool Add>
44  class CASADI_EXPORT SetNonzeros : public MXNode {
45  public:
47 
52  static MX create(const MX& y, const MX& x, const std::vector<casadi_int>& nz);
53  static MX create(const MX& y, const MX& x, const Slice& s);
54  static MX create(const MX& y, const MX& x, const Slice& inner, const Slice& outer);
56 
58  SetNonzeros(const MX& y, const MX& x);
59 
61  ~SetNonzeros() override = 0;
62 
64  virtual std::vector<casadi_int> all() const = 0;
65 
69  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
70  const std::vector<bool>& unique={}) const override;
71 
75  void eval_linear(const std::vector<std::array<MX, 3> >& arg,
76  std::vector<std::array<MX, 3> >& res) const override {
77  eval_linear_rearrange(arg, res);
78  }
79 
83  int eval_activity(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override {
84  return sp_forward(arg, res, iw, w);
85  }
86 
90  void ad_forward(const std::vector<std::vector<MX> >& fseed,
91  std::vector<std::vector<MX> >& fsens) const override;
92 
96  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
97  std::vector<std::vector<MX> >& asens) const override;
98 
102  casadi_int op() const override { return Add ? OP_ADDNONZEROS : OP_SETNONZEROS;}
103 
105  Matrix<casadi_int> mapping() const override;
106 
108  casadi_int n_inplace() const override { return 1;}
109 
113  static MXNode* deserialize(DeserializingStream& s);
114 
115  protected:
119  explicit SetNonzeros(DeserializingStream& s) : MXNode(s) {}
120  };
121 
122 
129  template<bool Add>
130  class CASADI_EXPORT SetNonzerosVector : public SetNonzeros<Add>{
131  public:
132 
134  SetNonzerosVector(const MX& y, const MX& x, const std::vector<casadi_int>& nz);
135 
137  ~SetNonzerosVector() override {}
138 
140  std::vector<casadi_int> all() const override { return nz_;}
141 
145  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
146  const std::vector<bool>& unique={}) const override;
147 
149  template<typename T>
150  int eval_gen(const T** arg, T** res, casadi_int* iw, T* w) const;
151 
153  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
154 
156  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
157 
161  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
162 
166  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
167 
171  std::string disp(const std::vector<std::string>& arg) const override;
172 
176  void generate(CodeGenerator& g,
177  const std::vector<casadi_int>& arg,
178  const std::vector<casadi_int>& res,
179  const std::vector<bool>& arg_is_ref,
180  std::vector<bool>& res_is_ref) const override;
181 
185  bool is_equal(const MXNode* node, casadi_int depth) const override;
186 
188  Dict info() const override { return {{"nz", nz_}, {"add", Add}}; }
189 
191  std::vector<casadi_int> nz_;
192 
196  void serialize_body(SerializingStream& s) const override;
200  void serialize_type(SerializingStream& s) const override;
201 
205  explicit SetNonzerosVector(DeserializingStream& s);
206  };
207 
208  // Specialization of the above when nz_ is a Slice
209  template<bool Add>
210  class CASADI_EXPORT SetNonzerosSlice : public SetNonzeros<Add>{
211  public:
212 
214  SetNonzerosSlice(const MX& y, const MX& x, const Slice& s) : SetNonzeros<Add>(y, x), s_(s) {}
215 
217  ~SetNonzerosSlice() override {}
218 
220  std::vector<casadi_int> all() const override { return s_.all(s_.stop);}
221 
225  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
226 
230  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
231 
235  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
236  const std::vector<bool>& unique={}) const override;
237 
239  template<typename T>
240  int eval_gen(const T** arg, T** res, casadi_int* iw, T* w) const;
241 
243  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
244 
246  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
247 
251  std::string disp(const std::vector<std::string>& arg) const override;
252 
256  void generate(CodeGenerator& g,
257  const std::vector<casadi_int>& arg,
258  const std::vector<casadi_int>& res,
259  const std::vector<bool>& arg_is_ref,
260  std::vector<bool>& res_is_ref) const override;
261 
265  bool is_equal(const MXNode* node, casadi_int depth) const override;
266 
268  Dict info() const override { return {{"slice", s_.info()}, {"add", Add}}; }
269 
270  // Data member
271  Slice s_;
272 
276  void serialize_body(SerializingStream& s) const override;
280  void serialize_type(SerializingStream& s) const override;
281 
285  explicit SetNonzerosSlice(DeserializingStream& s);
286  };
287 
288  // Specialization of the above when nz_ is a nested Slice
289  template<bool Add>
290  class CASADI_EXPORT SetNonzerosSlice2 : public SetNonzeros<Add>{
291  public:
292 
294  SetNonzerosSlice2(const MX& y, const MX& x, const Slice& inner, const Slice& outer) :
295  SetNonzeros<Add>(y, x), inner_(inner), outer_(outer) {}
296 
298  ~SetNonzerosSlice2() override {}
299 
301  std::vector<casadi_int> all() const override { return inner_.all(outer_, outer_.stop);}
302 
306  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
307 
311  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
312 
316  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
317  const std::vector<bool>& unique={}) const override;
318 
320  template<typename T>
321  int eval_gen(const T** arg, T** res, casadi_int* iw, T* w) const;
322 
324  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
325 
327  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
328 
332  std::string disp(const std::vector<std::string>& arg) const override;
333 
337  void generate(CodeGenerator& g,
338  const std::vector<casadi_int>& arg,
339  const std::vector<casadi_int>& res,
340  const std::vector<bool>& arg_is_ref,
341  std::vector<bool>& res_is_ref) const override;
342 
346  bool is_equal(const MXNode* node, casadi_int depth) const override;
347 
349  Dict info() const override { return {{"inner", inner_.info()}, {"outer", outer_.info()},
350  {"add", Add}}; }
351 
352  // Data members
353  Slice inner_, outer_;
354 
358  void serialize_body(SerializingStream& s) const override;
362  void serialize_type(SerializingStream& s) const override;
363 
367  explicit SetNonzerosSlice2(DeserializingStream& s);
368  };
369 
370 } // namespace casadi
372 
373 #endif // CASADI_SETNONZEROS_HPP
The casadi namespace.
Definition: archiver.hpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.