use crate::bindings::*;
use crate::{Error, Loss, Matrix, Model};
pub struct Params {
param: MfParameter,
}
impl Params {
pub(crate) fn new() -> Self {
let mut param = unsafe { mf_get_default_param() };
param.nr_bins = 25;
Self { param }
}
pub fn loss(&mut self, value: Loss) -> &mut Self {
self.param.fun = value;
self
}
pub fn factors(&mut self, value: i32) -> &mut Self {
self.param.k = value;
self
}
pub fn threads(&mut self, value: i32) -> &mut Self {
self.param.nr_threads = value;
self
}
pub fn bins(&mut self, value: i32) -> &mut Self {
self.param.nr_bins = value;
self
}
pub fn iterations(&mut self, value: i32) -> &mut Self {
self.param.nr_iters = value;
self
}
pub fn lambda_p1(&mut self, value: f32) -> &mut Self {
self.param.lambda_p1 = value;
self
}
pub fn lambda_p2(&mut self, value: f32) -> &mut Self {
self.param.lambda_p2 = value;
self
}
pub fn lambda_q1(&mut self, value: f32) -> &mut Self {
self.param.lambda_q1 = value;
self
}
pub fn lambda_q2(&mut self, value: f32) -> &mut Self {
self.param.lambda_q2 = value;
self
}
pub fn learning_rate(&mut self, value: f32) -> &mut Self {
self.param.eta = value;
self
}
pub fn alpha(&mut self, value: f32) -> &mut Self {
self.param.alpha = value;
self
}
pub fn c(&mut self, value: f32) -> &mut Self {
self.param.c = value;
self
}
pub fn nmf(&mut self, value: bool) -> &mut Self {
self.param.do_nmf = value;
self
}
pub fn quiet(&mut self, value: bool) -> &mut Self {
self.param.quiet = value;
self
}
pub fn fit(&self, data: &Matrix) -> Result<Model, Error> {
if data.is_empty() {
return Err(Error::Parameter("no data"));
}
let prob = data.try_into()?;
let param = self.build_param()?;
let model = unsafe { mf_train(&prob, param) };
if model.is_null() {
return Err(Error::Unknown);
}
Ok(Model { model })
}
pub fn fit_eval(&self, train_set: &Matrix, eval_set: &Matrix) -> Result<Model, Error> {
if train_set.is_empty() || eval_set.is_empty() {
return Err(Error::Parameter("no data"));
}
let tr: MfProblem = train_set.try_into()?;
let va: MfProblem = eval_set.try_into()?;
let param = self.build_param()?;
if matches!(param.fun, Loss::OneClassL2) {
if va.m > tr.m {
return Err(Error::Parameter(
"eval set cannot have extra rows for OneClassL2 loss",
));
}
if va.n > tr.n {
return Err(Error::Parameter(
"eval set cannot have extra columns for OneClassL2 loss",
));
}
}
let model = unsafe { mf_train_with_validation(&tr, &va, param) };
if model.is_null() {
return Err(Error::Unknown);
}
Ok(Model { model })
}
pub fn cv(&self, data: &Matrix, folds: i32) -> Result<f64, Error> {
if data.is_empty() {
return Err(Error::Parameter("no data"));
}
let prob = data.try_into()?;
let param = self.build_param()?;
let avg_error = unsafe { mf_cross_validation(&prob, folds, param) };
if avg_error == 0.0 {
return Err(Error::Unknown);
}
Ok(avg_error)
}
fn build_param(&self) -> Result<MfParameter, Error> {
let param = self.param;
if param.k < 1 {
return Err(Error::Parameter(
"number of factors must be greater than zero",
));
}
if param.nr_threads < 1 {
return Err(Error::Parameter(
"number of threads must be greater than zero",
));
}
if param.nr_bins < 1 || param.nr_bins < param.nr_threads {
return Err(Error::Parameter(
"number of bins must be greater than number of threads",
));
}
if param.nr_iters < 1 {
return Err(Error::Parameter(
"number of iterations must be greater than zero",
));
}
if param.lambda_p1 < 0.0
|| param.lambda_p2 < 0.0
|| param.lambda_q1 < 0.0
|| param.lambda_q2 < 0.0
{
return Err(Error::Parameter(
"regularization coefficient must be non-negative",
));
}
if param.eta <= 0.0 {
return Err(Error::Parameter("learning rate must be greater than zero"));
}
if matches!(param.fun, Loss::RealKL) && !param.do_nmf {
return Err(Error::Parameter(
"nmf must be set when using generalized KL-divergence",
));
}
if param.alpha < 0.0 {
return Err(Error::Parameter("alpha must be a non-negative number"));
}
Ok(param)
}
}