#![forbid(unsafe_code)]
mod error;
pub use error::{Error, Result};
mod lgb;
pub use lgb::LgbModel;
#[cfg(feature = "onnx")]
mod onnx;
#[cfg(feature = "onnx")]
pub use onnx::OnnxModel;
#[cfg(feature = "catboost")]
mod catboost;
#[cfg(feature = "catboost")]
pub use catboost::{CatBoostModel, Output};
use std::path::Path;
pub trait Model: Send + Sync {
fn predict(&self, features: &[f64]) -> Result<f64>;
fn predict_batch(
&self,
flat: &[f64],
n_features: usize,
) -> Result<Vec<f64>> {
check_batch_shape(flat, n_features)?;
flat.chunks_exact(n_features)
.map(|s| self.predict(s))
.collect()
}
}
fn check_batch_shape(flat: &[f64], n_features: usize) -> Result<()> {
if n_features == 0 || flat.len() % n_features != 0 {
return Err(Error::BatchShape {
len: flat.len(),
n_features,
});
}
Ok(())
}
pub fn load_model(path: &Path) -> Result<Box<dyn Model>> {
let ext = path.extension().and_then(|e| e.to_str()).unwrap_or("");
match ext.to_ascii_lowercase().as_str() {
"lgb" | "txt" => Ok(Box::new(LgbModel::load(path)?)),
#[cfg(feature = "onnx")]
"onnx" => Ok(Box::new(OnnxModel::load(path)?)),
#[cfg(feature = "catboost")]
"cbm" => Ok(Box::new(CatBoostModel::load(
path,
Output::Probability,
)?)),
_ => Err(Error::UnsupportedFormat {
ext: ext.to_string(),
supported: supported_extensions(),
}),
}
}
#[must_use]
pub fn supported_extensions() -> &'static [&'static str] {
&[
#[cfg(feature = "onnx")]
"onnx",
"lgb",
"txt",
#[cfg(feature = "catboost")]
"cbm",
]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_load_model_lgb_extensions() {
for name in ["bosk_case_test.LGB", "bosk_case_test.txt"] {
let path = std::env::temp_dir().join(name);
std::fs::copy("tests/fixtures/tiny_binary.lgb", &path)
.expect("copy fixture");
let model = load_model(&path).expect(name);
let _ = std::fs::remove_file(&path);
assert!(model
.predict(&[0.5, -0.2, 0.7, 1.1, -0.9])
.is_ok());
}
}
#[test]
fn test_supported_extensions_match_load_model() {
for ext in supported_extensions() {
let path = format!("no_such_model.{ext}");
let Err(err) = load_model(Path::new(&path)) else {
panic!(
"{ext}: expected an error for a nonexistent file"
);
};
assert!(
!matches!(err, Error::UnsupportedFormat { .. }),
"{ext} is listed but load_model refused it: {err:?}"
);
}
}
#[test]
fn test_load_model_unknown_extension() {
let Err(err) = load_model(Path::new("model.xgboost")) else {
panic!("expected UnsupportedFormat, got Ok");
};
assert!(
matches!(err, Error::UnsupportedFormat { .. }),
"got {err:?}"
);
}
}