ennbo-py 0.2.11

Python bindings for ENN core algorithms
//! Stateful ENN fitter Python bindings.

use numpy::{PyArray1, PyReadonlyArray2, ToPyArray};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use rand::rngs::StdRng;
use rand::SeedableRng;

use crate::py_model::{PyENNParams, PyEpistemicNearestNeighbors};

#[pyclass(name = "ENNStatefulFitter")]
pub struct PyENNStatefulFitter {
    inner: ennbo::ENNFitter,
    rng: StdRng,
}

#[pymethods]
impl PyENNStatefulFitter {
    #[new]
    #[pyo3(signature = (k, seed, infer_aleatoric_variance_scale=true))]
    #[doc = "kiss-coverage-off"]
    fn new(k: i32, seed: u64, infer_aleatoric_variance_scale: bool) -> Self {
        Self {
            inner: ennbo::ENNFitter::new(k, infer_aleatoric_variance_scale),
            rng: StdRng::seed_from_u64(seed),
        }
    }

    #[pyo3(signature = (x, y, yvar=None, y_bounds=None))]
    #[doc = "kiss-coverage-off"]
    fn tell(
        &mut self,
        x: PyReadonlyArray2<f64>,
        y: PyReadonlyArray2<f64>,
        yvar: Option<PyReadonlyArray2<f64>>,
        y_bounds: Option<PyReadonlyArray2<f64>>,
    ) -> PyResult<()> {
        let yvar_arr = yvar.as_ref().map(|v| v.as_array());
        let y_bounds_owned = y_bounds.as_ref().map(|v| v.as_array().to_owned());
        self.inner
            .tell(
                &x.as_array(),
                &y.as_array(),
                yvar_arr.as_ref(),
                y_bounds_owned.as_ref(),
            )
            .map_err(|e| PyValueError::new_err(e.to_string()))
    }

    #[doc = "kiss-coverage-off"]
    fn y_std<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyArray1<f64>>> {
        Ok(self.inner.y_std().to_pyarray_bound(py))
    }

    #[pyo3(signature = (model, num_fit_candidates, num_fit_samples, params_warm_start=None))]
    #[doc = "kiss-coverage-off"]
    fn ask(
        &mut self,
        model: &PyEpistemicNearestNeighbors,
        num_fit_candidates: usize,
        num_fit_samples: usize,
        params_warm_start: Option<PyENNParams>,
    ) -> PyResult<PyENNParams> {
        let warm = params_warm_start.as_ref().map(|p| p.inner);
        let result = self
            .inner
            .ask(
                &model.inner,
                num_fit_candidates,
                num_fit_samples,
                warm.as_ref(),
                &mut self.rng,
            )
            .map_err(|e| PyValueError::new_err(e.to_string()))?;
        Ok(PyENNParams { inner: result })
    }
}

#[cfg(test)]
mod kiss_coverage_tests {
    use super::*;

    #[test]
    fn py_fitter_units_are_linked() {
        let _ = (
            PyENNStatefulFitter::new,
            PyENNStatefulFitter::tell,
            PyENNStatefulFitter::y_std,
            PyENNStatefulFitter::ask,
            std::mem::size_of::<PyENNStatefulFitter>,
        );
    }
}