#![allow(dead_code, reason = "each test crate uses a subset of the helpers")]
use hessboost::model::Iterations;
use hessboost::prelude::{BoostedModel, DMatrix, HessboostError, Result};
pub fn labeled_dense(x: &[f32], n_cols: usize, labels: &[f32]) -> DMatrix {
DMatrix::from_dense(x, labels.len(), n_cols)
.unwrap()
.with_labels(labels)
.unwrap()
}
pub fn lcg(seed: u64) -> impl FnMut() -> f32 {
let mut s = seed;
move || {
s = s
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(s >> 40) as f32 / (1u32 << 24) as f32
}
}
pub fn four_features(i: usize) -> [f32; 4] {
let a = ((i * 37) % 101) as f32 / 101.0;
let b = ((i * 53) % 97) as f32 / 97.0;
let c = ((i * 11) % 89) as f32 / 89.0;
let d = if i.is_multiple_of(7) {
f32::NAN
} else {
((i * 29) % 83) as f32 / 83.0
};
[a, b, c, d]
}
pub fn invalid_param<T: std::fmt::Debug>(result: Result<T>) -> &'static str {
match result {
Err(HessboostError::InvalidParameter { name, .. }) => name,
other => panic!("expected an invalid-parameter error, got {other:?}"),
}
}
pub fn invalid_data<T: std::fmt::Debug>(result: Result<T>) -> (&'static str, Option<String>) {
match result {
Err(HessboostError::InvalidData { input, dataset, .. }) => (input, dataset),
other => panic!("expected an invalid-data error, got {other:?}"),
}
}
pub fn incompatible_model<T: std::fmt::Debug>(result: Result<T>) -> &'static str {
match result {
Err(HessboostError::IncompatibleModel { what, .. }) => what,
other => panic!("expected an incompatible-model error, got {other:?}"),
}
}
pub fn rmse(model: &BoostedModel, data: &DMatrix) -> f64 {
let preds = model.predict(data, Iterations::Best).unwrap();
let labels = data.labels().unwrap();
let sse: f64 = preds
.as_slice()
.iter()
.zip(labels)
.map(|(p, y)| f64::from(p - y).powi(2))
.sum();
(sse / labels.len() as f64).sqrt()
}
pub fn with_threads<T: Send>(threads: usize, f: impl FnOnce() -> T + Send) -> T {
rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap()
.install(f)
}
pub mod bits;
pub mod fixtures;
pub mod smooth;