use std::time::Instant;
use anyhow::{Result, bail};
use super::lipschitz::estimate_lipschitz;
use super::operator::QuadraticOperator;
use super::polish::improved_proxqp_like_polish_box_qp;
use super::problem::{BoxQPProblem, FirstOrderProblem};
use super::runtime_log::emit_log_line;
use super::types::{Diagnostics, SolverOptions, SolverQuality, SolverResult, SolverTiming};
use super::utils::{DenseVec, add_assign, copy_into, dot_diff, fill_momentum, fill_scaled_sub};
pub(crate) fn solve_box_qp_apgd<P: FirstOrderProblem>(
problem: &P,
operator: &dyn QuadraticOperator,
structured_problem: Option<&BoxQPProblem>,
options: &SolverOptions,
) -> Result<SolverResult> {
let total_start = Instant::now();
let n = problem.n();
if options.stopping.check_every == 0 {
bail!("stopping.check_every must be positive");
}
if options.logging.print_every == 0 {
bail!("logging.print_every must be positive");
}
if let Some(x0) = &options.x0 {
if x0.len() != n {
bail!("x0 has len {}, expected {}", x0.len(), n);
}
}
let mut x = options
.x0
.clone()
.map(DenseVec::from_vec)
.unwrap_or_else(|| DenseVec::zeros(n));
problem.clip_in_place(&mut x);
let l_value = options
.lipschitz
.value
.unwrap_or_else(|| estimate_lipschitz(operator, &options.lipschitz.method));
if !l_value.is_finite() || l_value <= 0.0 {
bail!("L must be finite and positive, got {l_value}");
}
let step = 1.0 / l_value;
let mut qx = DenseVec::zeros(n);
problem.matvec_into(&x, &mut qx);
let mut qx_old = qx.clone();
let mut y = x.clone();
let mut qy = qx.clone();
let mut x_old = x.clone();
let mut grad_y = DenseVec::zeros(n);
let mut x_candidate = DenseVec::zeros(n);
let mut t = 1.0;
let mut diag = problem.diagnostics_from_qx(
&qx,
&x,
options.stopping.bound_tol,
options.stopping.dual_certification,
);
let mut num_restarts = 0usize;
let mut iterations = 0usize;
let apgd_start = Instant::now();
if options.logging.verbose {
log_preamble(n, structured_problem.is_some(), options);
log_header(options.stopping.dual_certification);
}
if !converged(&diag, options.stopping.tol) {
for k in 1..=options.stopping.max_iter {
iterations = k;
copy_into(&mut x_old, &x);
copy_into(&mut qx_old, &qx);
copy_into(&mut grad_y, &qy);
add_assign(&mut grad_y, problem.c());
fill_scaled_sub(&mut x_candidate, &y, step, &grad_y);
problem.clip_in_place(&mut x_candidate);
if dot_diff(&x_candidate, &x_old, &y, &x_candidate) > 0.0 {
t = 1.0;
copy_into(&mut y, &x_old);
copy_into(&mut qy, &qx_old);
copy_into(&mut grad_y, &qy);
add_assign(&mut grad_y, problem.c());
fill_scaled_sub(&mut x_candidate, &y, step, &grad_y);
problem.clip_in_place(&mut x_candidate);
num_restarts += 1;
}
std::mem::swap(&mut x, &mut x_candidate);
problem.matvec_into(&x, &mut qx);
let t_new = 0.5_f64 * (1.0_f64 + (1.0_f64 + 4.0_f64 * t * t).sqrt());
let beta = (t - 1.0) / t_new;
fill_momentum(&mut y, &x, &x_old, beta);
fill_momentum(&mut qy, &qx, &qx_old, beta);
t = t_new;
if should_check_iteration(k, options.stopping.max_iter, options.stopping.check_every) {
let objective = problem.objective_from_qx(&x, &qx);
diag = problem.diagnostics_from_qx(
&qx,
&x,
options.stopping.bound_tol,
options.stopping.dual_certification,
);
if options.logging.verbose
&& (k == 1
|| k % options.logging.print_every == 0
|| converged(&diag, options.stopping.tol))
{
log_iteration(
k,
objective,
&diag,
num_restarts,
apgd_start.elapsed().as_secs_f64(),
);
}
if converged(&diag, options.stopping.tol) {
break;
}
}
}
}
let apgd_time = apgd_start.elapsed().as_secs_f64();
let mut objective = problem.objective_from_qx(&x, &qx);
let mut polish_time = 0.0;
if options.polish.enabled && should_run_polish(&diag, options) {
let Some(problem) = structured_problem else {
diag = problem.diagnostics_from_qx(
&qx,
&x,
options.stopping.bound_tol,
options.stopping.dual_certification,
);
return Ok(SolverResult {
x: x.iter().copied().collect(),
objective,
iterations,
num_restarts,
quality: quality_summary(&diag),
timing: SolverTiming {
apgd_time_sec: apgd_time,
polish_time_sec: 0.0,
total_time_sec: total_start.elapsed().as_secs_f64(),
},
lipschitz: l_value,
step_size: step,
scaling: super::types::ScalingSummary {
applied: false,
name: "none",
scale_min: 1.0,
scale_max: 1.0,
},
});
};
let before = objective;
let (x_polished, qx_polished, elapsed) =
improved_proxqp_like_polish_box_qp(problem, &x, &qx, options);
polish_time = elapsed;
let polished_objective = problem.objective_from_qx(&x_polished, &qx_polished);
if polished_objective <= before + 1e-10 * before.abs().max(1.0) {
x = x_polished;
qx = qx_polished;
objective = polished_objective;
}
}
diag = problem.diagnostics_from_qx(
&qx,
&x,
options.stopping.bound_tol,
options.stopping.dual_certification,
);
Ok(SolverResult {
x: x.iter().copied().collect(),
objective,
iterations,
num_restarts,
quality: quality_summary(&diag),
timing: SolverTiming {
apgd_time_sec: apgd_time,
polish_time_sec: polish_time,
total_time_sec: total_start.elapsed().as_secs_f64(),
},
lipschitz: l_value,
step_size: step,
scaling: super::types::ScalingSummary {
applied: false,
name: "none",
scale_min: 1.0,
scale_max: 1.0,
},
})
}
pub(crate) fn quality_summary(diag: &Diagnostics) -> SolverQuality {
SolverQuality {
gap: diag.gap,
rel_gap: diag.rel_gap,
certified_lower_bound: diag.certified_lower_bound,
kkt_inf: diag.kkt_inf,
}
}
fn converged(diag: &Diagnostics, tol: f64) -> bool {
if diag.rel_gap.is_finite() {
diag.rel_gap <= tol && diag.rel_kkt <= tol
} else {
diag.rel_kkt <= tol
}
}
fn should_run_polish(diag: &Diagnostics, options: &SolverOptions) -> bool {
if diag.rel_gap.is_finite() {
diag.rel_gap > 0.25 * options.stopping.tol
|| diag.rel_kkt > 0.25 * options.stopping.tol
} else {
diag.rel_kkt > 0.25 * options.stopping.tol
}
}
fn should_check_iteration(k: usize, max_iter: usize, check_every: usize) -> bool {
k == 1 || k % check_every == 0 || k == max_iter
}
fn log_preamble(n: usize, structured_problem: bool, options: &SolverOptions) {
let problem_kind = if structured_problem {
"structured box QP"
} else {
"implicit box QP"
};
let scaling = match options.scaling.mode {
super::types::ScalingMode::None => "none",
super::types::ScalingMode::HessianDiag => "hessian_diag",
};
let polish = if structured_problem && options.polish.enabled {
"enabled"
} else {
"disabled"
};
emit_log_line("");
emit_log_line("==== HerculesABQP solve ====");
emit_log_line(&format!("problem: {problem_kind}"));
emit_log_line(&format!("variables: {n}"));
emit_log_line(&format!("scaling: {scaling}"));
emit_log_line(&format!("dual certification: {}", options.stopping.dual_certification));
emit_log_line(&format!("polish: {polish}"));
emit_log_line(&format!("tolerance: {:.3e}", options.stopping.tol));
emit_log_line("");
}
fn log_header(dual_certification: bool) {
let header = if dual_certification {
format!(
"{:>8} {:>14} {:>10} {:>10} {:>10} {:>10} {:>9} {:>8}",
"iter", "objective", "gap", "rel_gap", "rel_kkt", "kkt_inf", "restarts", "time(s)"
)
} else {
format!(
"{:>8} {:>14} {:>10} {:>10} {:>9} {:>8}",
"iter", "objective", "rel_kkt", "kkt_inf", "restarts", "time(s)"
)
};
emit_log_line(&header);
emit_log_line(&"-".repeat(header.len()));
}
fn log_iteration(
k: usize,
objective: f64,
diag: &Diagnostics,
num_restarts: usize,
elapsed_sec: f64,
) {
if diag.rel_gap.is_finite() {
emit_log_line(&format!(
"{:8} {:14.6e} {:10.3e} {:10.3e} {:10.3e} {:10.3e} {:9} {:8.3}",
k,
objective,
diag.gap,
diag.rel_gap,
diag.rel_kkt,
diag.kkt_inf,
num_restarts,
elapsed_sec
));
} else {
emit_log_line(&format!(
"{:8} {:14.6e} {:10.3e} {:10.3e} {:9} {:8.3}",
k, objective, diag.rel_kkt, diag.kkt_inf, num_restarts, elapsed_sec
));
}
}