solve.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_SOLVE_HPP
27 #define CASADI_SOLVE_HPP
28 
29 #include "mx_node.hpp"
30 #include "casadi_call.hpp"
31 
32 namespace casadi {
46  template<bool Tr>
47  class CASADI_EXPORT Solve : public MXNode {
48  public:
52  Solve(const MX& r, const MX& A);
53 
57  ~Solve() override {}
58 
62  std::string disp(const std::vector<std::string>& arg) const override;
63 
67  virtual std::string mod_prefix() const {return "";}
68 
72  virtual std::string mod_suffix() const {return "";}
73 
77  void eval_mx(const std::vector<MX>& arg, std::vector<MX>& res,
78  const std::vector<bool>& unique={}) const override;
79 
83  void ad_forward(const std::vector<std::vector<MX> >& fseed,
84  std::vector<std::vector<MX> >& fsens) const override;
85 
89  void ad_reverse(const std::vector<std::vector<MX> >& aseed,
90  std::vector<std::vector<MX> >& asens) const override;
91 
93  casadi_int n_inplace() const override { return 1;}
94 
98  int sp_forward(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
99 
103  int eval_activity(const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
104 
108  int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w) const override;
109 
113  casadi_int op() const override { return OP_SOLVE;}
114 
116  Dict info() const override {
117  return {{"tr", Tr}};
118  }
119 
121  virtual MX solve(const MX& A, const MX& B, bool tr) const = 0;
122 
124  virtual const Sparsity& A_sp() const { return dep(1).sparsity();}
125 
129  void serialize_body(SerializingStream& s) const override;
130 
134  void serialize_type(SerializingStream& s) const override;
135 
139  static MXNode* deserialize(DeserializingStream& s);
140 
144  explicit Solve(DeserializingStream& s);
145  };
146 
153  template<bool Tr>
154  class CASADI_EXPORT LinsolCall : public Solve<Tr> {
155  public:
156 
160  LinsolCall(const MX& r, const MX& A, const Linsol& linear_solver);
161 
165  ~LinsolCall() override {}
166 
168  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
169 
171  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
172 
176  size_t sz_w() const override;
177 
181  size_t codegen_sz_w() const override;
182 
186  void generate(CodeGenerator& g,
187  const std::vector<casadi_int>& arg,
188  const std::vector<casadi_int>& res,
189  const std::vector<bool>& arg_is_ref,
190  std::vector<bool>& res_is_ref) const override;
191 
194 
196  MX solve(const MX& A, const MX& B, bool tr) const override {
197  return linsol_.solve(A, B, tr);
198  }
199 
203  void serialize_body(SerializingStream& s) const override;
204 
208  void serialize_type(SerializingStream& s) const override;
209 
213  static MXNode* deserialize(DeserializingStream& s);
214 
218  explicit LinsolCall(DeserializingStream& s);
219  };
220 
227  template<bool Tr>
228  class CASADI_EXPORT TriuSolve : public Solve<Tr> {
229  public:
230 
234  TriuSolve(const MX& r, const MX& A);
235 
239  ~TriuSolve() override {}
240 
242  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
243 
245  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
246 
248  MX solve(const MX& A, const MX& B, bool tr) const override {
249  return A->get_solve_triu(B, tr);
250  }
251 
255  void generate(CodeGenerator& g,
256  const std::vector<casadi_int>& arg,
257  const std::vector<casadi_int>& res,
258  const std::vector<bool>& arg_is_ref,
259  std::vector<bool>& res_is_ref) const override;
260  };
261 
268  template<bool Tr>
269  class CASADI_EXPORT TrilSolve : public Solve<Tr> {
270  public:
271 
275  TrilSolve(const MX& r, const MX& A);
276 
280  ~TrilSolve() override {}
281 
283  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
284 
286  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
287 
289  MX solve(const MX& A, const MX& B, bool tr) const override {
290  return A->get_solve_tril(B, tr);
291  }
292 
296  void generate(CodeGenerator& g,
297  const std::vector<casadi_int>& arg,
298  const std::vector<casadi_int>& res,
299  const std::vector<bool>& arg_is_ref,
300  std::vector<bool>& res_is_ref) const override;
301  };
302 
309  template<bool Tr>
310  class CASADI_EXPORT SolveUnity : public Solve<Tr> {
311  public:
312 
316  SolveUnity(const MX& r, const MX& A);
317 
321  ~SolveUnity() override {}
322 
326  std::string mod_prefix() const override {return "(I - ";}
327 
331  std::string mod_suffix() const override {return ")";}
332 
334  const Sparsity& A_sp() const override;
335 
336  // Sparsity pattern of linear system, cached
337  mutable Sparsity A_sp_;
338 
339 #ifdef CASADI_WITH_THREADSAFE_SYMBOLICS
341  mutable std::mutex A_sp_mtx_;
342 #endif // CASADI_WITH_THREADSAFE_SYMBOLICS
343  };
344 
351  template<bool Tr>
352  class CASADI_EXPORT TriuSolveUnity : public SolveUnity<Tr> {
353  public:
354 
358  TriuSolveUnity(const MX& r, const MX& A);
359 
363  ~TriuSolveUnity() override {}
364 
366  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
367 
369  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
370 
372  MX solve(const MX& A, const MX& B, bool tr) const override {
373  return A->get_solve_triu_unity(B, tr);
374  }
375 
379  void generate(CodeGenerator& g,
380  const std::vector<casadi_int>& arg,
381  const std::vector<casadi_int>& res,
382  const std::vector<bool>& arg_is_ref,
383  std::vector<bool>& res_is_ref) const override;
384  };
385 
392  template<bool Tr>
393  class CASADI_EXPORT TrilSolveUnity : public SolveUnity<Tr> {
394  public:
395 
399  TrilSolveUnity(const MX& r, const MX& A);
400 
404  ~TrilSolveUnity() override {}
405 
407  int eval(const double** arg, double** res, casadi_int* iw, double* w) const override;
408 
410  int eval_sx(const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w) const override;
411 
413  MX solve(const MX& A, const MX& B, bool tr) const override {
414  return A->get_solve_tril_unity(B, tr);
415  }
416 
420  void generate(CodeGenerator& g,
421  const std::vector<casadi_int>& arg,
422  const std::vector<casadi_int>& res,
423  const std::vector<bool>& arg_is_ref,
424  std::vector<bool>& res_is_ref) const override;
425  };
426 
427 } // namespace casadi
428 
429 #endif // CASADI_SOLVE_HPP
Helper class for C code generation.
Helper class for Serialization.
Linear solve operation with a linear solver instance.
Definition: solve.hpp:154
~LinsolCall() override
Destructor.
Definition: solve.hpp:165
MX solve(const MX &A, const MX &B, bool tr) const override
Solve another system with the same factorization.
Definition: solve.hpp:196
Linsol linsol_
Linear solver (may be shared between multiple nodes)
Definition: solve.hpp:193
Linear solver.
Definition: linsol.hpp:55
DM solve(const DM &A, const DM &B, bool tr=false) const
Node class for MX objects.
Definition: mx_node.hpp:51
MX - Matrix expression.
Definition: mx.hpp:92
Helper class for Serialization.
Linear solve with unity diagonal added.
Definition: solve.hpp:310
~SolveUnity() override
Destructor.
Definition: solve.hpp:321
Sparsity A_sp_
Definition: solve.hpp:337
std::string mod_suffix() const override
Modifier for linear system, after argument.
Definition: solve.hpp:331
std::string mod_prefix() const override
Modifier for linear system, before argument.
Definition: solve.hpp:326
An MX atomic for linear solver solution: x = r * A^-1 or x = r * A^-T.
Definition: solve.hpp:47
Dict info() const override
Definition: solve.hpp:116
virtual const Sparsity & A_sp() const
Sparsity pattern for the linear system.
Definition: solve.hpp:124
virtual std::string mod_prefix() const
Modifier for linear system, before argument.
Definition: solve.hpp:67
casadi_int op() const override
Get the operation.
Definition: solve.hpp:113
virtual std::string mod_suffix() const
Modifier for linear system, after argument.
Definition: solve.hpp:72
virtual MX solve(const MX &A, const MX &B, bool tr) const =0
Solve another system with the same factorization.
~Solve() override
Destructor.
Definition: solve.hpp:57
casadi_int n_inplace() const override
Can the operation be performed inplace (i.e. overwrite the result)
Definition: solve.hpp:93
General sparsity class.
Definition: sparsity.hpp:106
Linear solve with an upper triangular matrix.
Definition: solve.hpp:393
~TrilSolveUnity() override
Destructor.
Definition: solve.hpp:404
MX solve(const MX &A, const MX &B, bool tr) const override
Solve another system with the same factorization.
Definition: solve.hpp:413
Linear solve with an upper triangular matrix.
Definition: solve.hpp:269
MX solve(const MX &A, const MX &B, bool tr) const override
Solve another system with the same factorization.
Definition: solve.hpp:289
~TrilSolve() override
Destructor.
Definition: solve.hpp:280
Linear solve with an upper triangular matrix, unity diagonal.
Definition: solve.hpp:352
~TriuSolveUnity() override
Destructor.
Definition: solve.hpp:363
MX solve(const MX &A, const MX &B, bool tr) const override
Solve another system with the same factorization.
Definition: solve.hpp:372
Linear solve with an upper triangular matrix.
Definition: solve.hpp:228
~TriuSolve() override
Destructor.
Definition: solve.hpp:239
MX solve(const MX &A, const MX &B, bool tr) const override
Solve another system with the same factorization.
Definition: solve.hpp:248
The casadi namespace.
Definition: archiver.hpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.