use crate::error::{CatBoostError, CatBoostResult};
use crate::features::{
ObjectsOrderFeatures,
EmptyTextFeatures,
EmptyEmbeddingFeatures
};
use crate::sys;
use std::ffi::{CStr,CString};
use std::path::Path;
use std::os::raw::c_char;
pub struct Model {
handle: *mut sys::ModelCalcerHandle,
}
unsafe impl Send for Model {}
unsafe impl Sync for Model {}
impl Model {
fn new() -> Self {
let model_handle = unsafe { sys::ModelCalcerCreate() };
Model {
handle: model_handle,
}
}
pub fn load<P: AsRef<Path>>(path: P) -> CatBoostResult<Self> {
let model = Model::new();
let path_c_str = CString::new(path.as_ref().to_str().unwrap()).unwrap();
CatBoostError::check_return_value(unsafe {
sys::LoadFullModelFromFile(model.handle, path_c_str.as_ptr())
})?;
Ok(model)
}
pub fn load_buffer<P: AsRef<Vec<u8>>>(buffer: P) -> CatBoostResult<Self> {
let model = Model::new();
CatBoostError::check_return_value(unsafe {
sys::LoadFullModelFromBuffer(
model.handle,
buffer.as_ref().as_ptr() as *const std::os::raw::c_void,
buffer.as_ref().len(),
)
})?;
Ok(model)
}
fn set_or_check_object_count<
TFeature,
TObjectFeatures: AsRef<[TFeature]>,
TFeatures: AsRef<[TObjectFeatures]>
>
(
object_count: &mut Option<usize>,
features: &TFeatures
) -> CatBoostResult<()> {
let features_array_size = features.as_ref().len();
if features_array_size > 0 {
match object_count {
Some(count) => {
if *count != features_array_size {
return Err(
CatBoostError{ description: "features arguments have different nonzero sizes".to_owned() }
)
}
}
None => {
object_count.replace(features_array_size);
}
}
}
Ok(())
}
pub fn predict<
TObjectFloatFeatures: AsRef<[f32]>,
TFloatFeatures: AsRef<[TObjectFloatFeatures]>,
TCatFeatureString: AsRef<str>,
TObjectCatFeatures: AsRef<[TCatFeatureString]>,
TCatFeatures: AsRef<[TObjectCatFeatures]>,
TTextFeatureString: AsRef<CStr>,
TObjectTextFeatures: AsRef<[TTextFeatureString]>,
TTextFeatures: AsRef<[TObjectTextFeatures]>,
TEmbedding: AsRef<[f32]>,
TObjectEmbeddingFeatures: AsRef<[TEmbedding]>,
TEmbeddingFeatures: AsRef<[TObjectEmbeddingFeatures]>
>(
&self,
features: ObjectsOrderFeatures<
TFloatFeatures,
TCatFeatures,
TTextFeatures,
TEmbeddingFeatures
>
) -> CatBoostResult<Vec<f64>> {
let mut object_count = None;
Self::set_or_check_object_count(&mut object_count, &features.float_features)?;
Self::set_or_check_object_count(&mut object_count, &features.cat_features)?;
Self::set_or_check_object_count(&mut object_count, &features.text_features)?;
Self::set_or_check_object_count(&mut object_count, &features.embedding_features)?;
if object_count.is_none() {
return Err(
CatBoostError{ description: "all features arguments are empty".to_owned() }
);
}
let mut float_features_ptr = features.float_features
.as_ref()
.iter()
.map(|x| x.as_ref().as_ptr())
.collect::<Vec<_>>();
let hashed_cat_features = features.cat_features
.as_ref()
.iter()
.map(|doc_cat_features| {
doc_cat_features
.as_ref()
.iter()
.map(|cat_feature| unsafe {
sys::GetStringCatFeatureHash(
cat_feature.as_ref().as_ptr() as *const std::os::raw::c_char,
cat_feature.as_ref().len(),
)
})
.collect::<Vec<_>>()
})
.collect::<Vec<_>>();
let mut hashed_cat_features_ptr = hashed_cat_features
.iter()
.map(|x| x.as_ptr())
.collect::<Vec<_>>();
let mut text_features_ptr_storage = features.text_features
.as_ref()
.iter()
.map(
|object_text_features|
object_text_features.as_ref()
.iter()
.map(|text|
text.as_ref().as_ptr()
)
.collect::<Vec<_>>()
)
.collect::<Vec<_>>();
let mut text_features_ptr = text_features_ptr_storage
.iter_mut()
.map(|object_texts_ptrs: &mut Vec<*const c_char>| object_texts_ptrs.as_mut_ptr())
.collect::<Vec<_>>();
let mut embedding_dimensions = if !features.embedding_features.as_ref().is_empty() {
features.embedding_features.as_ref()[0].as_ref().iter()
.map(|x| x.as_ref().len())
.collect::<Vec<_>>()
} else {
vec![]
};
let mut embedding_features_ptr_storage = features.embedding_features
.as_ref()
.iter()
.map(
|object_embeddings|
object_embeddings.as_ref()
.iter()
.map(|embedding|
embedding.as_ref().as_ptr()
)
.collect::<Vec<_>>()
)
.collect::<Vec<_>>();
let mut embedding_features_ptr = embedding_features_ptr_storage
.iter_mut()
.map(|object_embeddings_ptrs: &mut Vec<*const f32>| object_embeddings_ptrs.as_mut_ptr())
.collect::<Vec<_>>();
let mut prediction = vec![0.0; object_count.unwrap() * self.get_dimensions_count()];
CatBoostError::check_return_value(unsafe {
sys::CalcModelPredictionWithHashedCatFeaturesAndTextAndEmbeddingFeatures(
self.handle,
object_count.unwrap(),
float_features_ptr.as_mut_ptr(),
if features.float_features.as_ref().is_empty() { 0 } else { features.float_features.as_ref()[0].as_ref().len() },
hashed_cat_features_ptr.as_mut_ptr(),
if features.cat_features.as_ref().is_empty() { 0 } else { features.cat_features.as_ref()[0].as_ref().len() },
text_features_ptr.as_mut_ptr(),
if features.text_features.as_ref().is_empty() { 0 } else { features.text_features.as_ref()[0].as_ref().len() },
embedding_features_ptr.as_mut_ptr(),
embedding_dimensions.as_mut_ptr(),
embedding_dimensions.len(),
prediction.as_mut_ptr(),
prediction.len(),
)
})?;
Ok(prediction)
}
pub fn calc_model_prediction<
TFloatFeature: AsRef<[f32]>,
TFloatFeatures: AsRef<[TFloatFeature]>,
TString: AsRef<str>,
TCatFeature: AsRef<[TString]>,
TCatFeatures: AsRef<[TCatFeature]>
>
(
&self,
float_features: TFloatFeatures,
cat_features: TCatFeatures,
) -> CatBoostResult<Vec<f64>> {
self.predict(
ObjectsOrderFeatures{
float_features,
cat_features,
text_features: EmptyTextFeatures{},
embedding_features: EmptyEmbeddingFeatures{}
}
)
}
pub fn get_float_features_count(&self) -> usize {
unsafe { sys::GetFloatFeaturesCount(self.handle) }
}
pub fn get_cat_features_count(&self) -> usize {
unsafe { sys::GetCatFeaturesCount(self.handle) }
}
pub fn get_text_features_count(&self) -> usize {
unsafe { sys::GetTextFeaturesCount(self.handle) }
}
pub fn get_embedding_features_count(&self) -> usize {
unsafe { sys::GetEmbeddingFeaturesCount(self.handle) }
}
pub fn get_tree_count(&self) -> usize {
unsafe { sys::GetTreeCount(self.handle)}
}
pub fn get_dimensions_count(&self) -> usize {
unsafe { sys::GetDimensionsCount(self.handle)}
}
pub fn enable_gpu_evaluation(&self) -> CatBoostResult<()> {
CatBoostError::check_return_value( unsafe { sys::EnableGPUEvaluation(self.handle, 0) } )
}
}
impl Drop for Model {
fn drop(&mut self) {
unsafe { sys::ModelCalcerDelete(self.handle) };
}
}