bisection.cpp
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 #include "bisection.hpp"
26 #include <cmath>
27 #include <algorithm>
28 #include <limits>
29 
30 namespace casadi {
31  extern "C" int CASADI_ROOTFINDER_BISECTION_EXPORT
32  casadi_register_rootfinder_bisection(Rootfinder::Plugin *plugin) {
33  plugin->creator = Bisection::creator;
34  plugin->name = "bisection";
35  plugin->doc = Bisection::meta_doc.c_str();
36  plugin->version = CASADI_VERSION;
37  plugin->options = &Bisection::options_;
38  plugin->deserialize = &Bisection::deserialize;
39  return 0;
40  }
41 
42  extern "C" void CASADI_ROOTFINDER_BISECTION_EXPORT casadi_load_rootfinder_bisection() {
44  }
45 
46  Bisection::Bisection(const std::string &name, const Function &f)
47  : Rootfinder(name, f) {
48  }
49 
50  Bisection::~Bisection() {
51  clear_mem();
52  }
53 
54  const Options Bisection::options_ = {{&Rootfinder::options_},
55  {
56  {"abstol", {OT_DOUBLE, "Stopping criterion tolerance on ||g||__inf)"}},
57  {"abstol_step", {OT_DOUBLE, "Stopping tolerance on bracket width"}},
58  {"max_iter",
59  {OT_INT, "Maximum number of Newton iterations to perform before returning."}},
60  {"lb", {OT_DOUBLE, "lower bound"}},
61  {"ub", {OT_DOUBLE, "upper bound"}},
62  {"search_step", {OT_DOUBLE, "Step size for bracket searching"}},
63  {"max_search", {OT_INT, "Maximum bracket search iterations"}},
64  }};
65 
66  void Bisection::init(const Dict &opts) {
67  Rootfinder::init(opts);
68 
69  max_iter_ = 100;
70  abstol_ = 1e-9;
71  abstol_step_ = 1e-9;
72  lb_ = -1e12;
73  ub_ = 1e12;
74  search_step_ = 1.0;
75  max_search_ = 100;
76 
77  for (auto &op : opts) {
78  if (op.first == "max_iter") {
79  max_iter_ = op.second;
80  } else if (op.first == "abstol") {
81  abstol_ = op.second;
82  } else if (op.first == "abstol_step") {
83  abstol_step_ = op.second;
84  } else if (op.first == "lb") {
85  lb_ = op.second;
86  } else if (op.first == "ub") {
87  ub_ = op.second;
88  } else if (op.first == "search_step") {
89  search_step_ = op.second;
90  } else if (op.first == "max_search") {
91  max_search_ = op.second;
92  }
93  }
94 
95  casadi_assert(oracle_.n_in() > 0,
96  "Bisection: the supplied f must have at least one input.");
97  casadi_assert(n_ == 1, "Bisection only supports scalar equations (n=1).");
98  casadi_assert(lb_ < ub_, "lb must be strictly less than ub.");
99  }
100 
101  int Bisection::init_mem(void *mem) const {
102  if (Rootfinder::init_mem(mem)) return 1;
103  auto m = static_cast<BisectionMemory *>(mem);
104  m->return_status = 0;
105  m->iter = 0;
106  m->search_iter = 0;
107  return 0;
108  }
109 
110  void Bisection::set_work(void *mem, const double **&arg, double **&res,
111  casadi_int *&iw, double *&w) const {
112  Rootfinder::set_work(mem, arg, res, iw, w);
113  }
114 
115  int Bisection::solve(void *mem) const {
116  auto m = static_cast<BisectionMemory *>(mem);
117 
118  double f_val = 0.0;
119 
120  auto eval_f = [&](double x) -> double {
121  for (casadi_int i = 0; i < n_in_; ++i)
122  m->arg[i] = m->iarg[i];
123 
124  m->arg[iin_] = &x;
125 
126  for (casadi_int i = 0; i < n_out_; ++i)
127  m->res[i] = nullptr;
128  m->res[iout_] = &f_val;
129 
130  if (oracle_(m->arg, m->res, m->iw, m->w, 0)) {
131  f_val = std::numeric_limits<double>::quiet_NaN();
132  }
133  return f_val;
134  };
135 
136  double x0 = m->iarg[iin_][0];
137  x0 = std::max(lb_, std::min(ub_, x0));
138 
139  double f0 = eval_f(x0);
140 
141  if (std::isnan(f0)) return finish(m, x0, f0, 0.0, -1, false, SOLVER_RET_UNKNOWN);
142  if (std::fabs(f0) < abstol_) return finish(m, x0, f0, 0, 2, true, SOLVER_RET_SUCCESS);
143 
144  bool bracketed = false;
145  double a = x0, b = x0;
146  double fa = f0, fb = f0;
147 
148  for (m->search_iter = 1; m->search_iter <= max_search_; ++m->search_iter) {
149  if (a > lb_) {
150  a = std::max(lb_, a - search_step_);
151  fa = eval_f(a);
152  }
153  if (b < ub_) {
154  b = std::min(ub_, b + search_step_);
155  fb = eval_f(b);
156  }
157 
158  if (std::isnan(fa) || std::isnan(fb))
159  return finish(m, a, fa, b - a, -1, false, SOLVER_RET_UNKNOWN);
160 
161  if (fa * fb <= 0.0) {
162  bracketed = true;
163  break;
164  }
165 
166  if (a == lb_ && b == ub_ && fa * fb > 0.0) {
167  break;
168  }
169  }
170 
171  if (!bracketed) {
172  return finish(m, x0, f0, b - a, -2, false, SOLVER_RET_UNKNOWN);
173  }
174 
175  if (fa == 0.0) return finish(m, a, 0.0, b - a, 2, true, SOLVER_RET_SUCCESS);
176  if (fb == 0.0) return finish(m, b, 0.0, b - a, 2, true, SOLVER_RET_SUCCESS);
177 
178  double mid = a, f_mid_val = fa;
179 
180  for (m->iter = 0; m->iter < max_iter_; ++m->iter) {
181  mid = 0.5 * (a + b);
182  f_mid_val = eval_f(mid);
183 
184  if (std::isnan(f_mid_val))
185  return finish(m, mid, f_mid_val, b - a, -1, false, SOLVER_RET_UNKNOWN);
186 
187  if (std::fabs(f_mid_val) < abstol_)
188  return finish(m, mid, f_mid_val, b - a, 2, true, SOLVER_RET_SUCCESS);
189 
190  if ((b - a) < abstol_step_)
191  return finish(m, mid, f_mid_val, b - a, 1, true, SOLVER_RET_SUCCESS);
192 
193  if (fa * f_mid_val < 0.0) {
194  b = mid;
195  fb = f_mid_val;
196  } else {
197  a = mid;
198  fa = f_mid_val;
199  }
200  }
201 
202  return finish(m, mid, f_mid_val, b - a, 0, false, SOLVER_RET_LIMITED);
203  }
204 
205  int Bisection::finish(BisectionMemory *m, double x_sol, double f_sol, double width,
206  int status, bool success, UnifiedReturnStatus urs) const {
207  casadi_copy(&x_sol, 1, m->ires[iout_]);
208  m->return_status = status;
209  m->f_mid = f_sol;
210  m->bracket_width = width;
211  m->success = success;
212  m->unified_return_status = urs;
213  return 0;
214  }
215 
216  std::string Bisection::status_str(int status) {
217  switch (status) {
218  case 0:
219  return "max_iteration_reached";
220  case 1:
221  return "converged_bracket";
222  case 2:
223  return "converged_abstol";
224  case -1:
225  return "nan_encountered";
226  case -2:
227  return "failed_to_bracket_root";
228  default:
229  return "unknown";
230  }
231  }
232 
233  Dict Bisection::get_stats(void *mem) const {
234  Dict stats = Rootfinder::get_stats(mem);
235  auto m = static_cast<BisectionMemory *>(mem);
236  stats["return_status"] = status_str(m->return_status);
237  stats["search_iter"] = m->search_iter; // number of bracket-search steps taken
238  stats["iter_count"] = m->iter;
239  stats["f_mid"] = m->f_mid;
240  stats["bracket_width"] = m->bracket_width;
241  return stats;
242  }
243 
244  void Bisection::serialize_body(SerializingStream &s) const {
245  Rootfinder::serialize_body(s);
246  s.version("Bisection", 1);
247  s.pack("Bisection::max_iter", max_iter_);
248  s.pack("Bisection::max_search", max_search_);
249  s.pack("Bisection::search_step", search_step_);
250  s.pack("Bisection::abstol", abstol_);
251  s.pack("Bisection::abstol_step", abstol_step_);
252  s.pack("Bisection::lb", lb_);
253  s.pack("Bisection::ub", ub_);
254  }
255 
256  Bisection::Bisection(DeserializingStream &s) : Rootfinder(s) {
257  s.version("Bisection", 1);
258  s.unpack("Bisection::max_iter", max_iter_);
259  s.unpack("Bisection::max_search", max_search_);
260  s.unpack("Bisection::search_step", search_step_);
261  s.unpack("Bisection::abstol", abstol_);
262  s.unpack("Bisection::abstol_step", abstol_step_);
263  s.unpack("Bisection::lb", lb_);
264  s.unpack("Bisection::ub", ub_);
265  }
266 
267 } // namespace casadi
static void registerPlugin(const Plugin &plugin, bool needs_lock=true)
Register an integrator in the factory.
The casadi namespace.
Definition: archiver.cpp:28
int CASADI_ROOTFINDER_BISECTION_EXPORT casadi_register_rootfinder_bisection(Rootfinder::Plugin *plugin)
Definition: bisection.cpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
void CASADI_ROOTFINDER_BISECTION_EXPORT casadi_load_rootfinder_bisection()
Definition: bisection.cpp:42