25 #include "bisection.hpp"
31 extern "C" int CASADI_ROOTFINDER_BISECTION_EXPORT
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;
46 Bisection::Bisection(
const std::string &name,
const Function &f)
47 : Rootfinder(name, f) {
50 Bisection::~Bisection() {
54 const Options Bisection::options_ = {{&Rootfinder::options_},
56 {
"abstol", {
OT_DOUBLE,
"Stopping criterion tolerance on ||g||__inf)"}},
57 {
"abstol_step", {
OT_DOUBLE,
"Stopping tolerance on bracket width"}},
59 {
OT_INT,
"Maximum number of Newton iterations to perform before returning."}},
62 {
"search_step", {
OT_DOUBLE,
"Step size for bracket searching"}},
63 {
"max_search", {
OT_INT,
"Maximum bracket search iterations"}},
66 void Bisection::init(
const Dict &opts) {
67 Rootfinder::init(opts);
77 for (
auto &op : opts) {
78 if (op.first ==
"max_iter") {
79 max_iter_ = op.second;
80 }
else if (op.first ==
"abstol") {
82 }
else if (op.first ==
"abstol_step") {
83 abstol_step_ = op.second;
84 }
else if (op.first ==
"lb") {
86 }
else if (op.first ==
"ub") {
88 }
else if (op.first ==
"search_step") {
89 search_step_ = op.second;
90 }
else if (op.first ==
"max_search") {
91 max_search_ = op.second;
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.");
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;
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);
115 int Bisection::solve(
void *mem)
const {
116 auto m =
static_cast<BisectionMemory *
>(mem);
120 auto eval_f = [&](
double x) ->
double {
121 for (casadi_int i = 0; i < n_in_; ++i)
122 m->arg[i] = m->iarg[i];
126 for (casadi_int i = 0; i < n_out_; ++i)
128 m->res[iout_] = &f_val;
130 if (oracle_(m->arg, m->res, m->iw, m->w, 0)) {
131 f_val = std::numeric_limits<double>::quiet_NaN();
136 double x0 = m->iarg[iin_][0];
137 x0 = std::max(lb_, std::min(ub_, x0));
139 double f0 = eval_f(x0);
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);
144 bool bracketed =
false;
145 double a = x0, b = x0;
146 double fa = f0, fb = f0;
148 for (m->search_iter = 1; m->search_iter <= max_search_; ++m->search_iter) {
150 a = std::max(lb_, a - search_step_);
154 b = std::min(ub_, b + search_step_);
158 if (std::isnan(fa) || std::isnan(fb))
159 return finish(m, a, fa, b - a, -1,
false, SOLVER_RET_UNKNOWN);
161 if (fa * fb <= 0.0) {
166 if (a == lb_ && b == ub_ && fa * fb > 0.0) {
172 return finish(m, x0, f0, b - a, -2,
false, SOLVER_RET_UNKNOWN);
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);
178 double mid = a, f_mid_val = fa;
180 for (m->iter = 0; m->iter < max_iter_; ++m->iter) {
182 f_mid_val = eval_f(mid);
184 if (std::isnan(f_mid_val))
185 return finish(m, mid, f_mid_val, b - a, -1,
false, SOLVER_RET_UNKNOWN);
187 if (std::fabs(f_mid_val) < abstol_)
188 return finish(m, mid, f_mid_val, b - a, 2,
true, SOLVER_RET_SUCCESS);
190 if ((b - a) < abstol_step_)
191 return finish(m, mid, f_mid_val, b - a, 1,
true, SOLVER_RET_SUCCESS);
193 if (fa * f_mid_val < 0.0) {
202 return finish(m, mid, f_mid_val, b - a, 0,
false, SOLVER_RET_LIMITED);
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;
210 m->bracket_width = width;
211 m->success = success;
212 m->unified_return_status = urs;
216 std::string Bisection::status_str(
int status) {
219 return "max_iteration_reached";
221 return "converged_bracket";
223 return "converged_abstol";
225 return "nan_encountered";
227 return "failed_to_bracket_root";
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;
238 stats[
"iter_count"] = m->iter;
239 stats[
"f_mid"] = m->f_mid;
240 stats[
"bracket_width"] = m->bracket_width;
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_);
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_);
static void registerPlugin(const Plugin &plugin, bool needs_lock=true)
Register an integrator in the factory.
int CASADI_ROOTFINDER_BISECTION_EXPORT casadi_register_rootfinder_bisection(Rootfinder::Plugin *plugin)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
void CASADI_ROOTFINDER_BISECTION_EXPORT casadi_load_rootfinder_bisection()