use ndarray::{Array1, Array2};
use numpy::{PyArray1, PyReadonlyArray1, PyReadonlyArray2};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyModule};
use sprs::CsMat;
use std::sync::Arc;
use crate::matrix::QuadraticMatrix;
use crate::solver::runtime_log::install_log_sink;
use crate::solver::{
LipschitzMethod, PreparedImplicitSolver, PreparedSolver, QuadraticOperator, ScalingMode,
SolverOptions, SolverResult,
};
fn parse_scaling_mode(name: &str) -> PyResult<ScalingMode> {
match name {
"none" => Ok(ScalingMode::None),
"hessian_diag" => Ok(ScalingMode::HessianDiag),
other => Err(PyValueError::new_err(format!(
"unknown scaling mode '{other}', expected 'none' or 'hessian_diag'"
))),
}
}
fn parse_lipschitz_method(name: &str) -> PyResult<LipschitzMethod> {
match name {
"gershgorin" => Ok(LipschitzMethod::Gershgorin),
"auto" => Ok(LipschitzMethod::Auto),
other => Err(PyValueError::new_err(format!(
"unknown lipschitz_method '{other}', expected 'gershgorin' or 'auto'"
))),
}
}
fn pyerr(err: anyhow::Error) -> PyErr {
PyRuntimeError::new_err(err.to_string())
}
struct ScopedPythonLogSink;
impl Drop for ScopedPythonLogSink {
fn drop(&mut self) {
install_log_sink(None);
}
}
fn maybe_install_python_log_sink(verbose: bool) -> Option<ScopedPythonLogSink> {
if !verbose {
return None;
}
let sink = Arc::new(|line: &str| {
Python::with_gil(|py| {
if let Ok(sys) = py.import("sys")
&& let Ok(stdout) = sys.getattr("stdout")
{
let _ = stdout.call_method1("write", (format!("{line}\n"),));
let _ = stdout.call_method0("flush");
}
});
});
install_log_sink(Some(sink));
Some(ScopedPythonLogSink)
}
fn solver_options(
assume_symmetric: bool,
scaling: &str,
lipschitz_method: &str,
lipschitz_value: Option<f64>,
x0: Option<Vec<f64>>,
max_iter: usize,
tol: f64,
dual_certification: bool,
check_every: usize,
bound_tol: f64,
polish: bool,
verbose: bool,
print_every: usize,
) -> PyResult<SolverOptions> {
let mut options = SolverOptions::default();
options.assume_symmetric = assume_symmetric;
options.scaling.mode = parse_scaling_mode(scaling)?;
options.lipschitz.method = parse_lipschitz_method(lipschitz_method)?;
options.lipschitz.value = lipschitz_value;
options.x0 = x0;
options.stopping.max_iter = max_iter;
options.stopping.tol = tol;
options.stopping.dual_certification = dual_certification;
options.stopping.check_every = check_every;
options.stopping.bound_tol = bound_tol;
options.polish.enabled = polish;
options.logging.verbose = verbose;
options.logging.print_every = print_every;
Ok(options)
}
fn result_to_pydict(py: Python<'_>, result: SolverResult) -> PyResult<PyObject> {
let out = PyDict::new(py);
out.set_item("x", result.x)?;
out.set_item("objective", result.objective)?;
out.set_item("iterations", result.iterations)?;
out.set_item("num_restarts", result.num_restarts)?;
out.set_item("rel_gap", result.quality.rel_gap)?;
out.set_item("kkt_inf", result.quality.kkt_inf)?;
out.set_item("lipschitz", result.lipschitz)?;
out.set_item("step_size", result.step_size)?;
out.set_item("scaling_applied", result.scaling.applied)?;
out.set_item("scaling_name", result.scaling.name)?;
out.set_item("scale_min", result.scaling.scale_min)?;
out.set_item("scale_max", result.scaling.scale_max)?;
out.set_item("apgd_time_sec", result.timing.apgd_time_sec)?;
out.set_item("polish_time_sec", result.timing.polish_time_sec)?;
out.set_item("total_time_sec", result.timing.total_time_sec)?;
Ok(out.into_any().unbind().into())
}
fn quadratic_from_dense(q: PyReadonlyArray2<'_, f64>) -> PyResult<QuadraticMatrix> {
let view = q.as_array();
Ok(QuadraticMatrix::dense(
Array2::from_shape_vec((view.nrows(), view.ncols()), view.iter().copied().collect())
.map_err(|err| PyValueError::new_err(err.to_string()))?,
))
}
fn quadratic_from_sparse(q: &Bound<'_, PyAny>) -> PyResult<QuadraticMatrix> {
let csr = q.call_method0("tocsr")?;
let csr = csr.call_method0("sorted_indices")?;
let shape: (usize, usize) = csr.getattr("shape")?.extract()?;
let indptr: Vec<usize> = csr.getattr("indptr")?.call_method0("tolist")?.extract()?;
let indices: Vec<usize> = csr.getattr("indices")?.call_method0("tolist")?.extract()?;
let data: Vec<f64> = csr.getattr("data")?.call_method0("tolist")?.extract()?;
let matrix = CsMat::new(shape, indptr, indices, data);
Ok(QuadraticMatrix::sparse(matrix))
}
fn quadratic_from_py(q: &Bound<'_, PyAny>) -> PyResult<QuadraticMatrix> {
if let Ok(dense) = q.extract::<PyReadonlyArray2<'_, f64>>() {
return quadratic_from_dense(dense);
}
if q.hasattr("tocsr")? {
return quadratic_from_sparse(q);
}
Err(PyValueError::new_err(
"Q must be either a numpy.ndarray or a scipy sparse matrix/array",
))
}
fn vec_from_py(name: &str, x: PyReadonlyArray1<'_, f64>) -> PyResult<Vec<f64>> {
x.as_slice()
.map(|slice| slice.to_vec())
.map_err(|_| PyValueError::new_err(format!("{name} must be a contiguous 1D float64 array")))
}
struct PythonQuadraticOperator {
obj: Py<PyAny>,
n: usize,
}
impl PythonQuadraticOperator {
fn new(obj: Py<PyAny>) -> PyResult<Self> {
Python::with_gil(|py| {
let bound = obj.bind(py);
let n: usize = bound
.getattr("n")
.map_err(|_| {
PyValueError::new_err("implicit operator must expose an integer 'n' attribute")
})?
.extract()
.map_err(|_| {
PyValueError::new_err("implicit operator attribute 'n' must be an integer")
})?;
Ok(Self { obj, n })
})
}
}
impl QuadraticOperator for PythonQuadraticOperator {
fn n(&self) -> usize {
self.n
}
fn matvec_into(&self, x: &Array1<f64>, out: &mut Array1<f64>) {
Python::with_gil(|py| {
let x_py = PyArray1::from_vec(py, x.to_vec());
let y_obj = self
.obj
.bind(py)
.call_method1("matvec", (x_py,))
.expect("python implicit operator matvec(x) failed");
let y = y_obj
.extract::<PyReadonlyArray1<'_, f64>>()
.expect("python implicit operator matvec(x) must return a contiguous 1D float64 numpy array");
let y = y
.as_slice()
.expect("python implicit operator matvec(x) must return a contiguous 1D float64 numpy array");
assert_eq!(
y.len(),
out.len(),
"python implicit operator matvec(x) returned len {}, expected {}",
y.len(),
out.len()
);
for (dst, src) in out.iter_mut().zip(y.iter().copied()) {
*dst = src;
}
});
}
fn diagonal(&self) -> Option<Array1<f64>> {
Python::with_gil(|py| {
let bound = self.obj.bind(py);
let diag_fn = bound.getattr("diagonal").ok()?;
let diag_obj = diag_fn.call0().ok()?;
let diag = diag_obj.extract::<PyReadonlyArray1<'_, f64>>().ok()?;
let diag = diag.as_slice().ok()?;
if diag.len() != self.n {
return None;
}
Some(Array1::from_vec(diag.to_vec()))
})
}
fn gershgorin_upper_bound(&self) -> Option<f64> {
Python::with_gil(|py| {
let bound = self.obj.bind(py);
let bound_fn = bound.getattr("gershgorin_upper_bound").ok()?;
bound_fn.call0().ok()?.extract::<f64>().ok()
})
}
}
#[pyfunction(
name = "solve_box_qp",
signature = (
q,
c,
lb,
ub,
*,
x0 = None,
assume_symmetric = false,
scaling = "hessian_diag",
lipschitz_method = "auto",
lipschitz_value = None,
max_iter = 5_000,
tol = 1e-6,
dual_certification = true,
check_every = 100,
bound_tol = 1e-10,
polish = true,
verbose = false,
print_every = 500
)
)]
fn py_solve_box_qp(
py: Python<'_>,
q: &Bound<'_, PyAny>,
c: PyReadonlyArray1<'_, f64>,
lb: PyReadonlyArray1<'_, f64>,
ub: PyReadonlyArray1<'_, f64>,
x0: Option<PyReadonlyArray1<'_, f64>>,
assume_symmetric: bool,
scaling: &str,
lipschitz_method: &str,
lipschitz_value: Option<f64>,
max_iter: usize,
tol: f64,
dual_certification: bool,
check_every: usize,
bound_tol: f64,
polish: bool,
verbose: bool,
print_every: usize,
) -> PyResult<PyObject> {
let _log_sink = maybe_install_python_log_sink(verbose);
let q = quadratic_from_py(q)?;
let c = vec_from_py("c", c)?;
let lb = vec_from_py("lb", lb)?;
let ub = vec_from_py("ub", ub)?;
let x0 = match x0 {
Some(x0) => Some(vec_from_py("x0", x0)?),
None => None,
};
let options = solver_options(
assume_symmetric,
scaling,
lipschitz_method,
lipschitz_value,
x0,
max_iter,
tol,
dual_certification,
check_every,
bound_tol,
polish,
verbose,
print_every,
)?;
let result = crate::solver::solve_box_qp(&q, &c, &lb, &ub, &options).map_err(pyerr)?;
result_to_pydict(py, result)
}
#[pyfunction(
name = "solve_box_qp_implicit",
signature = (
operator,
c,
lb,
ub,
*,
x0 = None,
assume_symmetric = true,
scaling = "none",
lipschitz_method = "auto",
lipschitz_value = None,
max_iter = 5_000,
tol = 1e-6,
dual_certification = true,
check_every = 100,
bound_tol = 1e-10,
verbose = false,
print_every = 500
)
)]
fn py_solve_box_qp_implicit(
py: Python<'_>,
operator: Py<PyAny>,
c: PyReadonlyArray1<'_, f64>,
lb: PyReadonlyArray1<'_, f64>,
ub: PyReadonlyArray1<'_, f64>,
x0: Option<PyReadonlyArray1<'_, f64>>,
assume_symmetric: bool,
scaling: &str,
lipschitz_method: &str,
lipschitz_value: Option<f64>,
max_iter: usize,
tol: f64,
dual_certification: bool,
check_every: usize,
bound_tol: f64,
verbose: bool,
print_every: usize,
) -> PyResult<PyObject> {
let _log_sink = maybe_install_python_log_sink(verbose);
let operator = Arc::new(PythonQuadraticOperator::new(operator)?);
let c = vec_from_py("c", c)?;
let lb = vec_from_py("lb", lb)?;
let ub = vec_from_py("ub", ub)?;
let x0 = match x0 {
Some(x0) => Some(vec_from_py("x0", x0)?),
None => None,
};
let options = solver_options(
assume_symmetric,
scaling,
lipschitz_method,
lipschitz_value,
x0,
max_iter,
tol,
dual_certification,
check_every,
bound_tol,
false,
verbose,
print_every,
)?;
let result =
crate::solver::solve_box_qp_implicit(operator, &c, &lb, &ub, &options).map_err(pyerr)?;
result_to_pydict(py, result)
}
#[pyclass(name = "PreparedSolver")]
struct PyPreparedSolver {
inner: PreparedSolver,
base_options: SolverOptions,
}
#[pyclass(name = "PreparedImplicitSolver")]
struct PyPreparedImplicitSolver {
inner: PreparedImplicitSolver,
base_options: SolverOptions,
}
#[pymethods]
impl PyPreparedSolver {
#[new]
#[pyo3(signature = (
q,
c,
*,
assume_symmetric = false,
scaling = "hessian_diag",
lipschitz_method = "auto",
lipschitz_value = None
))]
fn new(
q: &Bound<'_, PyAny>,
c: PyReadonlyArray1<'_, f64>,
assume_symmetric: bool,
scaling: &str,
lipschitz_method: &str,
lipschitz_value: Option<f64>,
) -> PyResult<Self> {
let q = quadratic_from_py(q)?;
let c = vec_from_py("c", c)?;
let options = solver_options(
assume_symmetric,
scaling,
lipschitz_method,
lipschitz_value,
None,
5_000,
1e-6,
true,
100,
1e-10,
true,
false,
500,
)?;
let inner = PreparedSolver::new(&q, &c, &options).map_err(pyerr)?;
Ok(Self {
inner,
base_options: options,
})
}
#[pyo3(signature = (
lb,
ub,
*,
x0 = None,
max_iter = 5_000,
tol = 1e-6,
dual_certification = true,
check_every = 100,
bound_tol = 1e-10,
polish = true,
verbose = false,
print_every = 500
))]
fn solve(
&self,
py: Python<'_>,
lb: PyReadonlyArray1<'_, f64>,
ub: PyReadonlyArray1<'_, f64>,
x0: Option<PyReadonlyArray1<'_, f64>>,
max_iter: usize,
tol: f64,
dual_certification: bool,
check_every: usize,
bound_tol: f64,
polish: bool,
verbose: bool,
print_every: usize,
) -> PyResult<PyObject> {
let _log_sink = maybe_install_python_log_sink(verbose);
let lb = vec_from_py("lb", lb)?;
let ub = vec_from_py("ub", ub)?;
let mut options = self.base_options.clone();
options.x0 = match x0 {
Some(x0) => Some(vec_from_py("x0", x0)?),
None => None,
};
options.stopping.max_iter = max_iter;
options.stopping.tol = tol;
options.stopping.dual_certification = dual_certification;
options.stopping.check_every = check_every;
options.stopping.bound_tol = bound_tol;
options.polish.enabled = polish;
options.logging.verbose = verbose;
options.logging.print_every = print_every;
let result = self.inner.solve(&lb, &ub, &options).map_err(pyerr)?;
result_to_pydict(py, result)
}
}
#[pymethods]
impl PyPreparedImplicitSolver {
#[new]
#[pyo3(signature = (
operator,
c,
*,
assume_symmetric = true,
scaling = "none",
lipschitz_method = "auto",
lipschitz_value = None
))]
fn new(
operator: Py<PyAny>,
c: PyReadonlyArray1<'_, f64>,
assume_symmetric: bool,
scaling: &str,
lipschitz_method: &str,
lipschitz_value: Option<f64>,
) -> PyResult<Self> {
let operator = Arc::new(PythonQuadraticOperator::new(operator)?);
let c = vec_from_py("c", c)?;
let options = solver_options(
assume_symmetric,
scaling,
lipschitz_method,
lipschitz_value,
None,
5_000,
1e-6,
true,
100,
1e-10,
false,
false,
500,
)?;
let inner = PreparedImplicitSolver::new(operator, &c, &options).map_err(pyerr)?;
Ok(Self {
inner,
base_options: options,
})
}
#[pyo3(signature = (
lb,
ub,
*,
x0 = None,
max_iter = 5_000,
tol = 1e-6,
dual_certification = true,
check_every = 100,
bound_tol = 1e-10,
verbose = false,
print_every = 500
))]
fn solve(
&self,
py: Python<'_>,
lb: PyReadonlyArray1<'_, f64>,
ub: PyReadonlyArray1<'_, f64>,
x0: Option<PyReadonlyArray1<'_, f64>>,
max_iter: usize,
tol: f64,
dual_certification: bool,
check_every: usize,
bound_tol: f64,
verbose: bool,
print_every: usize,
) -> PyResult<PyObject> {
let _log_sink = maybe_install_python_log_sink(verbose);
let lb = vec_from_py("lb", lb)?;
let ub = vec_from_py("ub", ub)?;
let mut options = self.base_options.clone();
options.x0 = match x0 {
Some(x0) => Some(vec_from_py("x0", x0)?),
None => None,
};
options.stopping.max_iter = max_iter;
options.stopping.tol = tol;
options.stopping.dual_certification = dual_certification;
options.stopping.check_every = check_every;
options.stopping.bound_tol = bound_tol;
options.polish.enabled = false;
options.logging.verbose = verbose;
options.logging.print_every = print_every;
let result = self.inner.solve(&lb, &ub, &options).map_err(pyerr)?;
result_to_pydict(py, result)
}
}
#[pymodule]
pub fn herculesabqp(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add(
"__doc__",
"Python bindings for HerculesABQP.\n\n\
The module exposes:\n\
- solve_box_qp(...) for explicit dense or sparse quadratic matrices\n\
- solve_box_qp_implicit(...) for Python-defined matrix-free operators\n\
- PreparedSolver for repeated explicit solves\n\
- PreparedImplicitSolver for repeated matrix-free solves\n\n\
All solver entry points return a dictionary with the primal solution,\n\
objective value, timing information, and compact convergence diagnostics.",
)?;
m.add_function(wrap_pyfunction!(py_solve_box_qp, m)?)?;
m.add_function(wrap_pyfunction!(py_solve_box_qp_implicit, m)?)?;
m.add_class::<PyPreparedSolver>()?;
m.add_class::<PyPreparedImplicitSolver>()?;
Ok(())
}