use numpy::*;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use crate::Label;
use crate::estimates::*;
use crate::fbleau_estimation::{Logger,run_fbleau};
#[pymodule(fbleau)]
fn pyfbleau(_py: Python, m: &PyModule) -> PyResult<()> {
#[pyfn(m, "run_fbleau")]
fn run_fbleau_py(py: Python,
train_x: &PyArray2<f64>, train_y: &PyArray1<Label>,
test_x: &PyArray2<f64>, test_y: &PyArray1<Label>,
estimate: &str, knn_strategy: Option<&str>,
distance: Option<String>,
log_errors: bool,
log_individual_errors: bool,
delta: Option<f64>, qstop: Option<usize>, absolute: bool,
scale: bool) -> PyResult<PyObject> {
let train_x = unsafe { train_x.as_array().to_owned() };
let train_y = unsafe { train_y.as_array().to_owned() };
let test_x = unsafe { test_x.as_array().to_owned() };
let test_y = unsafe { test_y.as_array().to_owned() };
let estimate = match estimate {
"nn" => Estimate::NN,
"knn" => Estimate::KNN,
"frequentist" => Estimate::Frequentist,
"nn-bound" => Estimate::NNBound,
_ => { unimplemented!() },
};
let knn_strategy = if let Some(strategy) = knn_strategy {
match strategy {
"ln" => Some(KNNStrategy::Ln),
"log10" => Some(KNNStrategy::Log10),
_ => { unimplemented!() },
}
} else {
None
};
let mut error_logger = if log_errors {
Some(Logger::LogVec(vec![]))
} else {
None
};
let mut individual_error_logger = if log_individual_errors {
Some(Logger::LogVec(vec![]))
} else {
None
};
let (min_error, last_error, random_guessing) =
run_fbleau(train_x, train_y, test_x, test_y, estimate,
knn_strategy, distance, &mut error_logger,
&mut individual_error_logger, delta, qstop, absolute,
scale);
let res = PyDict::new(py);
res.set_item("min-estimate", min_error)?;
res.set_item("last-estimate", last_error)?;
res.set_item("random-guessing", random_guessing)?;
res.set_item("estimates",
if let Some(Logger::LogVec(v)) = error_logger {
v
} else {
vec![]
})?;
res.set_item("min-individual-errors",
if let Some(Logger::LogVec(v)) = individual_error_logger {
v
} else {
vec![]
})?;
Ok(res.into())
}
Ok(())
}