41namespace ZonoOpt::detail {
46 struct ADMM_data : std::enable_shared_from_this<ADMM_data>
48 Eigen::SparseMatrix<zono_float> P, A, AT;
49 Eigen::SparseMatrix<zono_float, Eigen::RowMajor> A_rm;
51 Eigen::Vector<zono_float, 1> c;
52 LDLT_data ldlt_data_M, ldlt_data_AAT;
55 std::shared_ptr<Box> x_box;
59 ADMM_data() =
default;
61 ADMM_data(
const Eigen::SparseMatrix<zono_float>& P,
const Eigen::Vector<zono_float, -1>& q,
62 const Eigen::SparseMatrix<zono_float>& A,
const Eigen::Vector<zono_float, -1>& b,
63 const Eigen::Vector<zono_float, -1>& x_l,
const Eigen::Vector<zono_float, -1>& x_u,
66 set(P, q, A, b, x_l, x_u, c, settings);
70 void set(
const Eigen::SparseMatrix<zono_float>& P,
const Eigen::Vector<zono_float, -1>& q,
71 const Eigen::SparseMatrix<zono_float>& A,
const Eigen::Vector<zono_float, -1>& b,
72 const Eigen::Vector<zono_float, -1>& x_l,
const Eigen::Vector<zono_float, -1>& x_u,
78 this->AT = A.transpose();
83 this->n_x =
static_cast<int>(P.rows());
84 this->n_cons =
static_cast<int>(A.rows());
85 this->sqrt_n_x = std::sqrt(
static_cast<zono_float>(this->n_x));
87 this->x_box = std::make_shared<Box>(x_l, x_u);
89 if (!settings.
settings_valid())
throw std::invalid_argument(
"ADMM data: invalid settings.");
90 this->settings = settings;
94 ADMM_data* clone()
const
96 const auto new_data =
new ADMM_data(*
this);
97 new_data->x_box = std::make_shared<Box>(*this->x_box);
103 inline void print_str(std::stringstream &ss)
108 std::cout << ss.str() << std::endl;
128 explicit ADMM_solver(
const ADMM_data& data)
131 this->data = std::make_shared<ADMM_data>(data);
132 this->eps_prim = data.settings.eps_prim;
133 this->eps_dual = data.settings.eps_dual;
136 this->is_warmstarted =
false;
144 explicit ADMM_solver(
const std::shared_ptr<ADMM_data>& data)
148 this->eps_prim = data->settings.eps_prim;
149 this->eps_dual = data->settings.eps_dual;
152 this->is_warmstarted =
false;
160 ADMM_solver(
const ADMM_solver& other)
162 this->data = other.data;
165 this->is_warmstarted = other.is_warmstarted;
166 this->eps_dual = other.eps_dual;
167 this->eps_prim = other.eps_prim;
174 virtual ~ADMM_solver() =
default;
182 virtual void warmstart(
const Eigen::Vector<zono_float, -1>& x0,
183 const Eigen::Vector<zono_float, -1>& u0)
190 this->is_warmstarted =
true;
197 virtual void factorize()
199 auto t0 = std::chrono::high_resolution_clock::now();
201 std::stringstream ss;
202 if (!this->data->ldlt_data_M.factorized)
204 t0 = std::chrono::high_resolution_clock::now();
206 if (this->data->settings.verbose)
208 run_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
209 std::chrono::high_resolution_clock::now() - t0).count());
210 ss <<
"M factorization time = " << run_time <<
" sec";
214 if (!this->data->ldlt_data_AAT.factorized)
216 t0 = std::chrono::high_resolution_clock::now();
217 this->factorize_AAT();
218 if (this->data->settings.verbose)
220 run_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
221 std::chrono::high_resolution_clock::now() - t0).count());
222 ss <<
"A*A^T factorization time = " << run_time <<
" sec";
240 if (
const bool contractor_feasible = this->startup(*this->data->x_box, solution); !contractor_feasible)
246 solve_core(*this->data->x_box, solution, stop);
249 OptSolution solve() {
return this->solve(
nullptr); }
254 std::shared_ptr<ADMM_data> data;
258 bool startup(
Box& x_box,
OptSolution& solution,
const std::set<int>& contract_inds=std::set<int>())
261 const auto start = std::chrono::high_resolution_clock::now();
264 std::stringstream ss;
267 if (!this->check_problem_dimensions())
269 throw std::invalid_argument(
"ADMM solve: inconsistent problem data dimensions.");
271 if (this->data->settings.verbose)
273 ss <<
"Solving ADMM problem with " << this->data->n_x <<
" variables and " << this->data->n_cons <<
" constraints.";
281 bool contractor_feasible =
true;
282 if (this->data->settings.use_interval_contractor)
284 const auto t0 = std::chrono::high_resolution_clock::now();
285 if (contract_inds.empty())
287 contractor_feasible = x_box.
contract(this->data->A_rm, this->data->b, this->data->settings.contractor_iter);
291 contractor_feasible = x_box.
contract_subset(this->data->A_rm, this->data->b, this->data->settings.contractor_iter,
292 this->data->A, contract_inds, this->data->settings.contractor_tree_search_depth);
295 if (this->data->settings.verbose)
297 const double run_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
298 std::chrono::high_resolution_clock::now() - t0).count());
299 ss <<
"Interval contractor time = " << run_time <<
" sec";
305 const double startup_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
306 std::chrono::high_resolution_clock::now() - start).count());
311 if (!contractor_feasible)
313 if (this->data->settings.verbose)
315 ss <<
"Infeasibility detected via interval contractor";
322 solution.
primal_residual = std::numeric_limits<zono_float>::infinity();
323 solution.
dual_residual = std::numeric_limits<zono_float>::infinity();
324 solution.
J = std::numeric_limits<zono_float>::infinity();
325 solution.
x = Eigen::Vector<zono_float, -1>::Zero(this->data->n_x);
326 solution.
z = solution.
x;
327 solution.
u = solution.
x;
329 return contractor_feasible;
333 virtual void solve_core(
const Box& x_box,
OptSolution& solution, std::atomic<bool>* stop)
336 auto start = std::chrono::high_resolution_clock::now();
337 std::stringstream ss;
340 Eigen::Vector<
zono_float, -1> xk, zk, uk, zkm1, rhs, x_nu;
341 if (this->is_warmstarted)
349 uk = Eigen::Vector<zono_float, -1>::Zero(this->data->n_x);
352 rhs = Eigen::Vector<zono_float, -1>::Zero(this->data->n_x + this->data->n_cons);
353 rhs.segment(this->data->n_x, this->data->n_cons) = this->data->b;
361 double run_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
362 std::chrono::high_resolution_clock::now() - start).count());
363 bool converged =
false, infeasible =
false;
365 while ((k < this->data->settings.k_max_admm) && (run_time+solution.
startup_time < this->data->settings.t_max) && !converged && !infeasible
366 && !(stop && (*stop)))
369 rhs.segment(0, this->data->n_x) = -this->data->q + this->data->settings.rho*(zk - uk);
370 x_nu = solve_LDLT(this->data->ldlt_data_M, rhs);
371 xk = x_nu.segment(0, this->data->n_x);
381 if (k % this->data->settings.k_inf_check == 0)
383 infeasible = this->is_infeasibility_certificate(zk - xk, xk, x_box);
384 if (this->data->settings.verbose && infeasible)
386 ss <<
"Infeasibility certificate detected at iteration " << k;
392 if (this->data->settings.inf_norm_conv)
394 rp_k = (xk - zk).cwiseAbs().maxCoeff();
395 rd_k = this->data->settings.rho*(zk - zkm1).cwiseAbs().maxCoeff();
396 converged = (rp_k < this->eps_prim && rd_k < this->eps_dual);
400 rp_k = (xk - zk).norm();
401 rd_k = this->data->settings.rho*(zk - zkm1).norm();
402 converged = (rp_k < this->data->sqrt_n_x*this->eps_prim && rd_k < this->data->sqrt_n_x*this->eps_dual);
410 run_time = 1e-6 *
static_cast<double>(std::chrono::duration_cast<std::chrono::microseconds>(
411 std::chrono::high_resolution_clock::now() - start).count());
414 if (this->data->settings.verbose && (k % this->data->settings.verbosity_interval == 0))
416 ss <<
"k = " << k <<
": primal residual = " << rp_k <<
", dual residual = "
417 << rd_k <<
", run time = " << run_time <<
" sec";
423 if (this->data->settings.verbose)
427 ss <<
"ADMM converged in " << k <<
" iterations.";
432 ss <<
"ADMM detected infeasibility.";
437 ss <<
"ADMM did not converge in " << k <<
" iterations.";
443 this->is_warmstarted =
false;
449 solution.
J = (0.5*zk.transpose()*this->data->P*zk + this->data->q.transpose()*zk + this->data->c)(0);
462 bool is_warmstarted =
false;
465 void factorize_M()
const
468 Eigen::SparseMatrix<zono_float> M (this->data->n_x + this->data->n_cons, this->data->n_x + this->data->n_cons);
470 Eigen::SparseMatrix<zono_float> I (this->data->n_x, this->data->n_x);
472 Eigen::SparseMatrix<zono_float> Phi = this->data->P + this->data->settings.rho*I;
474 std::vector<Eigen::Triplet<zono_float>> triplets;
475 get_triplets_offset<zono_float>(Phi, triplets, 0, 0);
476 get_triplets_offset<zono_float>(this->data->A, triplets, this->data->n_x, 0);
477 get_triplets_offset<zono_float>(this->data->AT, triplets, 0, this->data->n_x);
478 M.setFromTriplets(triplets.begin(), triplets.end());
481 Eigen::SimplicialLDLT<Eigen::SparseMatrix<zono_float>> ldlt_solver_M;
482 ldlt_solver_M.compute(M);
483 if (ldlt_solver_M.info() != Eigen::Success)
484 throw std::runtime_error(
"ADMM: factorization of problem data failed, most likely A is not full row rank");
486 get_LDLT_data(ldlt_solver_M, this->data->ldlt_data_M);
489 void factorize_AAT()
const
492 const Eigen::SparseMatrix<zono_float> AAT = this->data->A*this->data->AT;
493 Eigen::SimplicialLDLT<Eigen::SparseMatrix<zono_float>> ldlt_solver_AAT;
494 ldlt_solver_AAT.compute(AAT);
495 if (ldlt_solver_AAT.info() != Eigen::Success)
496 throw std::runtime_error(
"ADMM: factorization of A*A^T failed, most likely A is not full row rank");
497 get_LDLT_data(ldlt_solver_AAT, this->data->ldlt_data_AAT);
501 bool is_infeasibility_certificate(
const Eigen::Vector<zono_float, -1>& ek,
502 const Eigen::Vector<zono_float, -1>& xk,
const Box& x_box)
const
505 Eigen::Vector<
zono_float, -1> A_e = this->data->A*ek;
506 Eigen::Vector<
zono_float, -1> AAT_inv_A_e = solve_LDLT(this->data->ldlt_data_AAT, A_e);
507 Eigen::Vector<
zono_float, -1> ek_proj = this->data->AT*AAT_inv_A_e;
515 bool check_problem_dimensions()
const
517 const bool prob_data_consistent = (this->data->P.rows() == this->data->n_x && this->data->P.cols() == this->data->n_x &&
518 this->data->q.size() == this->data->n_x && this->data->A.rows() == this->data->n_cons &&
519 this->data->A.cols() == this->data->n_x && this->data->b.size() == this->data->n_cons &&
520 this->data->x_box->size() == this->data->n_x);
522 bool warm_start_consistent;
523 if (this->is_warmstarted)
525 warm_start_consistent = (this->x0.size() == this->data->n_x && this->u0.size() == this->data->n_x &&
526 this->data->x_box->size() == this->data->n_x);
529 warm_start_consistent =
true;
531 return prob_data_consistent && warm_start_consistent;