26 #ifndef CASADI_BSPLINE_HPP
27 #define CASADI_BSPLINE_HPP
29 #include "mx_node.hpp"
42 class CASADI_EXPORT BSplineCommon :
public MXNode {
46 BSplineCommon(
const std::vector<double>& knots,
47 const std::vector<casadi_int>& offset,
48 const std::vector<casadi_int>& degree,
50 const std::vector<casadi_int>& lookup_mode);
53 ~BSplineCommon()
override {}
55 static void prepare(casadi_int m,
const std::vector<casadi_int>& offset,
56 const std::vector<casadi_int>& degree, casadi_int &coeffs_size,
57 std::vector<casadi_int>& coeffs_dims, std::vector<casadi_int>& strides);
59 static casadi_int get_coeff_size(casadi_int m,
const std::vector<casadi_int>& offset,
60 const std::vector<casadi_int>& degree);
63 static M derivative_coeff(casadi_int i,
64 const std::vector< std::vector<double> >& knots,
65 const std::vector<casadi_int>& degree,
66 const std::vector<casadi_int>& coeffs_dims,
68 std::vector< std::vector<double> >& new_knots,
69 std::vector<casadi_int>& new_degree);
71 std::vector<double> knots_;
72 std::vector<casadi_int> offset_;
73 std::vector<casadi_int> degree_;
75 std::vector<casadi_int> lookup_mode_;
78 std::vector<casadi_int> strides_;
79 std::vector<casadi_int> coeffs_dims_;
80 casadi_int coeffs_size_;
88 mutable MX jac_cache_;
90 #ifdef CASADI_WITH_THREADSAFE_SYMBOLICS
92 mutable std::mutex jac_cache_mtx_;
95 virtual MX jac_cached()
const = 0;
100 static size_t n_iw(
const std::vector<casadi_int> °ree);
105 static size_t n_w(
const std::vector<casadi_int> °ree);
110 size_t sz_iw()
const override;
115 size_t sz_w()
const override;
120 casadi_int op()
const override {
return OP_BSPLINE;}
125 void ad_forward(
const std::vector<std::vector<MX> >& fseed,
126 std::vector<std::vector<MX> >& fsens)
const override;
131 void ad_reverse(
const std::vector<std::vector<MX> >& aseed,
132 std::vector<std::vector<MX> >& asens)
const override;
137 void generate(CodeGenerator& g,
138 const std::vector<casadi_int>& arg,
139 const std::vector<casadi_int>& res,
140 const std::vector<bool>& arg_is_ref,
141 std::vector<bool>& res_is_ref)
const override;
146 virtual std::string generate(CodeGenerator& g,
147 const std::vector<casadi_int>& arg,
148 const std::vector<bool>& arg_is_ref)
const = 0;
153 static MXNode* deserialize(DeserializingStream& s);
156 MX jac(
const MX& x,
const T& coeffs)
const;
161 void serialize_body(SerializingStream& s)
const override;
168 explicit BSplineCommon(DeserializingStream& s);
180 class CASADI_EXPORT BSpline :
public BSplineCommon {
183 static MX create(
const MX& x,
const std::vector< std::vector<double> >& knots,
184 const std::vector<double>& coeffs,
185 const std::vector<casadi_int>& degree,
190 BSpline(
const MX& x,
const std::vector<double>& knots,
191 const std::vector<casadi_int>& offset,
192 const std::vector<double>& coeffs,
193 const std::vector<casadi_int>& degree,
195 const std::vector<casadi_int>& lookup_mode);
198 ~BSpline()
override {}
201 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override;
206 void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
207 const std::vector<bool>& unique={})
const override;
212 std::string generate(CodeGenerator& g,
213 const std::vector<casadi_int>& arg,
214 const std::vector<bool>& arg_is_ref)
const override;
219 std::string disp(
const std::vector<std::string>& arg)
const override;
222 std::vector<double> coeffs_;
224 MX jac_cached()
const override;
236 static DM dual(
const std::vector<double>& x,
237 const std::vector< std::vector<double> >& knots,
238 const std::vector<casadi_int>& degree,
243 void serialize_body(SerializingStream& s)
const override;
247 void serialize_type(SerializingStream& s)
const override;
252 explicit BSpline(DeserializingStream& s);
256 class CASADI_EXPORT BSplineParametric :
public BSplineCommon {
258 static MX create(
const MX& x,
const MX& coeffs,
259 const std::vector< std::vector<double> >& knots,
260 const std::vector<casadi_int>& degree,
265 static MX create(
const MX& x,
const MX& coeffs,
266 const std::vector<MX>& knots,
267 const std::vector<casadi_int>& degree,
272 BSplineParametric(
const MX& x,
const MX& coeffs,
273 const std::vector<double>& knots,
274 const std::vector<casadi_int>& offset,
275 const std::vector<casadi_int>& degree,
277 const std::vector<casadi_int>& lookup_mode);
280 ~BSplineParametric()
override {}
283 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override;
288 void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
289 const std::vector<bool>& unique={})
const override;
291 MX jac_cached()
const override;
296 std::string generate(CodeGenerator& g,
297 const std::vector<casadi_int>& arg,
298 const std::vector<bool>& arg_is_ref)
const override;
303 std::string disp(
const std::vector<std::string>& arg)
const override;
308 void serialize_type(SerializingStream& s)
const override;
313 explicit BSplineParametric(DeserializingStream& s) : BSplineCommon(s) {}
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.