split.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_SPLIT_HPP
27 #define CASADI_SPLIT_HPP
28 
29 #include "multiple_output.hpp"
30 #include <map>
31 #include <stack>
32 
34 
35 namespace casadi {
36 
43  class CASADI_EXPORT Split : public MultipleOutput {
44  public:
46  Split(const MX& x, const std::vector<casadi_int>& offset);
47 
49  ~Split() override = 0;
50 
54  casadi_int nout() const override { return output_sparsity_.size(); }
55 
59  const Sparsity& sparsity(casadi_int oind) const override { return output_sparsity_.at(oind);}
60 
62  template<typename T>
63  int eval_gen(const T** arg, T** res, casadi_int* iw, T* w) const;
64 
66  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
67 
69  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
70 
74  void eval_linear(const std::vector<std::array<MX, 3> >& arg,
75  std::vector<std::array<MX, 3> >& res) const override {
76  eval_linear_rearrange(arg, res);
77  }
78 
82  int eval_activity(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override {
83  return sp_forward(arg, res, iw, w);
84  }
85 
89  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
90 
94  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
95 
99  void generate(CodeGenerator& g,
100  const std::vector<casadi_int>& arg,
101  const std::vector<casadi_int>& res,
102  const std::vector<bool>& arg_is_ref,
103  std::vector<bool>& res_is_ref) const override;
104 
106  Dict info() const override;
107 
108  // Sparsity pattern of the outputs
109  std::vector<casadi_int> offset_;
110  std::vector<Sparsity> output_sparsity_;
111 
115  void serialize_body(SerializingStream& s) const override;
116 
117  protected:
121  explicit Split(DeserializingStream& s);
122  };
123 
130  class CASADI_EXPORT Horzsplit : public Split {
131  public:
132 
134  Horzsplit(const MX& x, const std::vector<casadi_int>& offset);
135 
137  ~Horzsplit() override {}
138 
142  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
143  const std::vector<bool>& unique={}) const override;
144 
148  void ad_forward(const std::vector<std::vector<MX> >& fseed,
149  std::vector<std::vector<MX> >& fsens) const override;
150 
154  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
155  std::vector<std::vector<MX> >& asens) const override;
156 
160  std::string disp(const std::vector<std::string>& arg) const override;
161 
165  casadi_int op() const override { return OP_HORZSPLIT;}
166 
168  MX get_horzcat(const std::vector<MX>& x) const override;
169 
173  static MXNode* deserialize(DeserializingStream& s) { return new Horzsplit(s); }
174 
175  protected:
179  explicit Horzsplit(DeserializingStream& s) : Split(s) {}
180  };
181 
188  class CASADI_EXPORT Diagsplit : public Split {
189  public:
190 
192  Diagsplit(const MX& x,
193  const std::vector<casadi_int>& offset1, const std::vector<casadi_int>& offset2);
194 
196  ~Diagsplit() override {}
197 
201  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
202  const std::vector<bool>& unique={}) const override;
203 
207  void ad_forward(const std::vector<std::vector<MX> >& fseed,
208  std::vector<std::vector<MX> >& fsens) const override;
209 
213  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
214  std::vector<std::vector<MX> >& asens) const override;
215 
219  std::string disp(const std::vector<std::string>& arg) const override;
220 
224  casadi_int op() const override { return OP_DIAGSPLIT;}
225 
227  MX get_diagcat(const std::vector<MX>& x) const override;
228 
232  static MXNode* deserialize(DeserializingStream& s) { return new Diagsplit(s); }
233 
234  protected:
238  explicit Diagsplit(DeserializingStream& s) : Split(s) {}
239  };
240 
247  class CASADI_EXPORT Vertsplit : public Split {
248  public:
249 
251  Vertsplit(const MX& x, const std::vector<casadi_int>& offset);
252 
254  ~Vertsplit() override {}
255 
259  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
260  const std::vector<bool>& unique={}) const override;
261 
265  void ad_forward(const std::vector<std::vector<MX> >& fseed,
266  std::vector<std::vector<MX> >& fsens) const override;
267 
271  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
272  std::vector<std::vector<MX> >& asens) const override;
273 
277  std::string disp(const std::vector<std::string>& arg) const override;
278 
282  casadi_int op() const override { return OP_VERTSPLIT;}
283 
285  MX get_vertcat(const std::vector<MX>& x) const override;
286 
290  static MXNode* deserialize(DeserializingStream& s) { return new Vertsplit(s); }
291 
292  protected:
296  explicit Vertsplit(DeserializingStream& s) : Split(s) {}
297  };
298 
299 } // namespace casadi
300 
302 
303 #endif // CASADI_SPLIT_HPP
The casadi namespace.
Definition: archiver.hpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.