use crate::rng::Rng;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use crate::check;
use crate::data::dmatrix::check_len;
use crate::data::{DMatrix, FeatureType};
use crate::error::{HessboostError, Result};
const NO_CATEGORY: u32 = u32::MAX;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum TargetKind {
#[default]
Regression,
Binary,
}
#[derive(Debug, Clone, PartialEq)]
pub struct OrderedTargetEncoder {
prior_weight: f64,
prior: Option<f64>,
permutations: usize,
seed: u64,
target: TargetKind,
}
#[derive(Debug, Clone, PartialEq)]
pub struct OrderedTargetEncoderBuilder {
encoder: OrderedTargetEncoder,
}
impl Default for OrderedTargetEncoderBuilder {
fn default() -> Self {
OrderedTargetEncoderBuilder {
encoder: OrderedTargetEncoder {
prior_weight: 1.0,
prior: None,
permutations: 1,
seed: 0,
target: TargetKind::Regression,
},
}
}
}
impl OrderedTargetEncoderBuilder {
#[must_use]
pub fn prior_weight(mut self, weight: f64) -> Self {
self.encoder.prior_weight = weight;
self
}
#[must_use]
pub fn prior(mut self, prior: f64) -> Self {
self.encoder.prior = Some(prior);
self
}
#[must_use]
pub fn permutations(mut self, permutations: usize) -> Self {
self.encoder.permutations = permutations;
self
}
#[must_use]
pub fn seed(mut self, seed: u64) -> Self {
self.encoder.seed = seed;
self
}
#[must_use]
pub fn target(mut self, target: TargetKind) -> Self {
self.encoder.target = target;
self
}
pub fn build(self) -> Result<OrderedTargetEncoder> {
let e = &self.encoder;
check::positive("prior_weight", e.prior_weight)?;
if let Some(p) = e.prior {
check::ensure(
"prior",
p.is_finite() && p.abs() <= f64::from(f32::MAX),
format!("must be finite and within the f32 range, got {p}"),
)?;
}
check::ensure("permutations", e.permutations != 0, "must be >= 1, got 0")?;
Ok(self.encoder)
}
}
struct ColumnCodes {
categories: Vec<u32>,
ids: Vec<u32>,
}
impl OrderedTargetEncoder {
pub fn builder() -> OrderedTargetEncoderBuilder {
OrderedTargetEncoderBuilder::default()
}
pub fn fit_transform(
&self,
data: &DMatrix,
columns: &[usize],
) -> Result<(DMatrix, FittedTargetEncoder)> {
let labels = data.labels().ok_or_else(|| {
HessboostError::invalid_data("labels", "ordered target statistics need labels")
})?;
if labels.len() != data.n_rows() {
return Err(HessboostError::invalid_data(
"labels",
"ordered target statistics need exactly one label per row; pass the per-row \
target separately",
));
}
self.fit_labels(data, columns, labels)
}
pub fn fit_transform_with_labels(
&self,
data: &DMatrix,
columns: &[usize],
labels: &[f32],
) -> Result<(DMatrix, FittedTargetEncoder)> {
check_len("target statistics labels", labels.len(), data.n_rows())?;
if labels.iter().any(|y| !y.is_finite()) {
return Err(HessboostError::invalid_data(
"labels",
"ordered target statistics need finite labels",
));
}
self.fit_labels(data, columns, labels)
}
fn fit_labels(
&self,
data: &DMatrix,
columns: &[usize],
labels: &[f32],
) -> Result<(DMatrix, FittedTargetEncoder)> {
let n_rows = data.n_rows();
if self.target == TargetKind::Binary && labels.iter().any(|&y| y != 0.0 && y != 1.0) {
return Err(HessboostError::invalid_data(
"labels",
"TargetKind::Binary needs 0/1 labels (multiclass targets are not supported)",
));
}
let slots = column_slots(data, columns)?;
let prior = self
.prior
.unwrap_or_else(|| labels.iter().map(|&y| f64::from(y)).sum::<f64>() / n_rows as f64);
let a = self.prior_weight;
let codes = collect_codes(data, &slots, columns.len());
let mut encodings = vec![vec![0f64; n_rows]; codes.len()];
let mut rng = Rng::new(self.seed);
let mut order: Vec<usize> = (0..n_rows).collect();
for _ in 0..self.permutations {
rng.shuffle(&mut order);
encodings
.par_iter_mut()
.zip(codes.par_iter())
.for_each(|(enc, col)| ordered_pass(&order, col, labels, prior, a, enc));
}
let scale = 1.0 / self.permutations as f64;
let overflow = || {
HessboostError::invalid_param("prior", "target statistics round outside the f32 range")
};
if encodings
.iter()
.flatten()
.any(|&v| !((v * scale) as f32).is_finite())
{
return Err(overflow());
}
let encoded = data
.map_values(|row, col, v| match slots[col] {
Some(s) => (encodings[s][row] * scale) as f32,
None => v,
})
.with_feature_types(&numeric_types(data, columns.iter().copied()))?;
let columns: Vec<ColumnEncoding> = codes
.into_iter()
.zip(columns)
.map(|(col, &column)| ColumnEncoding::fit(column, col, labels, prior, a))
.collect();
if columns
.iter()
.any(|c| c.values.iter().any(|v| !v.is_finite()))
{
return Err(overflow());
}
let fitted = FittedTargetEncoder {
n_cols: data.n_cols(),
prior: prior as f32,
columns,
};
Ok((encoded, fitted))
}
}
fn column_slots(data: &DMatrix, columns: &[usize]) -> Result<Vec<Option<usize>>> {
if columns.is_empty() {
return Err(HessboostError::invalid_param(
"columns",
"name at least one categorical column to encode",
));
}
let mut slots = vec![None; data.n_cols()];
for (slot, &col) in columns.iter().enumerate() {
if col >= data.n_cols() {
return Err(HessboostError::FeatureOutOfBounds {
index: col,
num_features: data.n_cols(),
});
}
if slots[col].replace(slot).is_some() {
return Err(HessboostError::invalid_param(
"columns",
format!("column {col} is listed twice"),
));
}
require_categorical(data, col)?;
}
Ok(slots)
}
fn require_categorical(data: &DMatrix, col: usize) -> Result<()> {
if data.feature_types()[col] == FeatureType::Categorical {
Ok(())
} else {
Err(HessboostError::invalid_param(
"columns",
format!("column {col} is not categorical; mark it with with_feature_types"),
))
}
}
fn numeric_types(data: &DMatrix, columns: impl IntoIterator<Item = usize>) -> Vec<FeatureType> {
let mut types = data.feature_types().to_vec();
for col in columns {
types[col] = FeatureType::Numerical;
}
types
}
fn collect_codes(data: &DMatrix, slots: &[Option<usize>], n_encoded: usize) -> Vec<ColumnCodes> {
let mut raw = vec![vec![NO_CATEGORY; data.n_rows()]; n_encoded];
data.for_each_entry(|row, col, v| {
if let Some(s) = slots[col as usize] {
raw[s][row] = v as u32;
}
});
raw.into_par_iter()
.map(|mut ids| {
let mut categories: Vec<u32> =
ids.iter().copied().filter(|&c| c != NO_CATEGORY).collect();
categories.sort_unstable();
categories.dedup();
for id in &mut ids {
if *id != NO_CATEGORY {
*id = categories.binary_search(id).unwrap_or_default() as u32;
}
}
ColumnCodes { categories, ids }
})
.collect()
}
fn ordered_pass(
order: &[usize],
col: &ColumnCodes,
labels: &[f32],
prior: f64,
a: f64,
enc: &mut [f64],
) {
let mut sums = vec![0f64; col.categories.len()];
let mut counts = vec![0f64; col.categories.len()];
for &row in order {
let id = col.ids[row];
if id == NO_CATEGORY {
continue;
}
let id = id as usize;
enc[row] += smoothed_mean(sums[id], counts[id], prior, a);
sums[id] += f64::from(labels[row]);
counts[id] += 1.0;
}
}
fn smoothed_mean(sum: f64, count: f64, prior: f64, a: f64) -> f64 {
let denom = count + a;
sum / denom + prior * (a / denom)
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
struct ColumnEncoding {
column: usize,
categories: Vec<u32>,
values: Vec<f32>,
}
impl ColumnEncoding {
fn fit(column: usize, codes: ColumnCodes, labels: &[f32], prior: f64, a: f64) -> Self {
let mut sums = vec![0f64; codes.categories.len()];
let mut counts = vec![0f64; codes.categories.len()];
for (&id, &y) in codes.ids.iter().zip(labels) {
if id != NO_CATEGORY {
sums[id as usize] += f64::from(y);
counts[id as usize] += 1.0;
}
}
let values = sums
.iter()
.zip(&counts)
.map(|(&sum, &count)| smoothed_mean(sum, count, prior, a) as f32)
.collect();
ColumnEncoding {
column,
categories: codes.categories,
values,
}
}
fn encode(&self, category: u32, prior: f32) -> f32 {
match self.categories.binary_search(&category) {
Ok(i) => self.values[i],
Err(_) => prior,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "FittedRepr")]
pub struct FittedTargetEncoder {
n_cols: usize,
prior: f32,
columns: Vec<ColumnEncoding>,
}
#[derive(Deserialize)]
struct FittedRepr {
n_cols: usize,
prior: f32,
columns: Vec<ColumnEncoding>,
}
impl TryFrom<FittedRepr> for FittedTargetEncoder {
type Error = HessboostError;
fn try_from(r: FittedRepr) -> Result<Self> {
let bad = |reason: &str| HessboostError::ModelFormat(format!("target encoder: {reason}"));
if !r.prior.is_finite() {
return Err(bad("prior must be finite"));
}
if r.columns.is_empty() {
return Err(bad("no encoded columns"));
}
let mut seen = std::collections::HashSet::with_capacity(r.columns.len());
for c in &r.columns {
if c.column >= r.n_cols || !seen.insert(c.column) {
return Err(bad("encoded columns must be distinct and < n_cols"));
}
if c.values.len() != c.categories.len() {
return Err(bad("categories and values differ in length"));
}
if c.categories.windows(2).any(|w| w[0] >= w[1]) {
return Err(bad("categories must be strictly ascending"));
}
if c.values.iter().any(|v| !v.is_finite()) {
return Err(bad("encoded values must be finite"));
}
}
Ok(FittedTargetEncoder {
n_cols: r.n_cols,
prior: r.prior,
columns: r.columns,
})
}
}
impl FittedTargetEncoder {
pub fn transform(&self, data: &DMatrix) -> Result<DMatrix> {
check_len("target encoder feature count", data.n_cols(), self.n_cols)?;
let mut slots = vec![None; self.n_cols];
for (s, c) in self.columns.iter().enumerate() {
require_categorical(data, c.column)?;
slots[c.column] = Some(s);
}
data.map_values(|_, col, v| match slots[col] {
Some(s) => self.columns[s].encode(v as u32, self.prior),
None => v,
})
.with_feature_types(&numeric_types(data, self.columns()))
}
pub fn encode(&self, column: usize, category: u32) -> Option<f32> {
self.columns
.iter()
.find(|c| c.column == column)
.map(|c| c.encode(category, self.prior))
}
pub fn prior(&self) -> f32 {
self.prior
}
pub fn n_cols(&self) -> usize {
self.n_cols
}
pub fn columns(&self) -> impl ExactSizeIterator<Item = usize> + '_ {
self.columns.iter().map(|c| c.column)
}
}
#[cfg(test)]
mod tests {
use super::*;
const CAT_NUM: [FeatureType; 2] = [FeatureType::Categorical, FeatureType::Numerical];
fn matrix(cats: &[f32], labels: &[f32]) -> DMatrix {
let x: Vec<f32> = cats
.iter()
.enumerate()
.flat_map(|(i, &c)| [c, i as f32])
.collect();
crate::test_support::labeled_dense(&x, cats.len(), 2, labels)
.with_feature_types(&CAT_NUM)
.unwrap()
}
fn column(m: &DMatrix, col: usize) -> Vec<Option<f32>> {
(0..m.n_rows()).map(|r| m.get(r, col)).collect()
}
fn encoder(permutations: usize, seed: u64) -> OrderedTargetEncoder {
OrderedTargetEncoder::builder()
.prior(0.5)
.permutations(permutations)
.seed(seed)
.build()
.unwrap()
}
fn crafted() -> (Vec<f32>, Vec<f32>) {
let cats: Vec<f32> = (0..40).map(|i| (i % 3) as f32).collect();
let labels: Vec<f32> = (0..40).map(|i| ((i * 7 + 3) % 5) as f32 * 0.5).collect();
(cats, labels)
}
#[test]
fn training_encoding_never_uses_own_target() {
let (cats, labels) = crafted();
for permutations in [1, 3] {
let enc = encoder(permutations, 11);
let (base, _) = enc.fit_transform(&matrix(&cats, &labels), &[0]).unwrap();
let base = column(&base, 0);
let mut others_moved = false;
for i in 0..labels.len() {
let mut perturbed = labels.clone();
perturbed[i] += 7.0;
let (out, _) = enc.fit_transform(&matrix(&cats, &perturbed), &[0]).unwrap();
let out = column(&out, 0);
assert_eq!(
out[i].map(f32::to_bits),
base[i].map(f32::to_bits),
"row {i}"
);
others_moved |= out != base;
}
assert!(others_moved);
}
}
#[test]
fn single_permutation_follows_the_ordered_formula() {
let (a, p) = (1.5, 0.25);
let n = 12;
let enc = OrderedTargetEncoder::builder()
.prior(p)
.prior_weight(a)
.seed(3)
.build()
.unwrap();
let (out, _) = enc
.fit_transform(&matrix(&vec![4.0; n], &vec![2.0; n]), &[0])
.unwrap();
let mut got: Vec<f32> = column(&out, 0).into_iter().map(Option::unwrap).collect();
got.sort_by(f32::total_cmp);
let want: Vec<f32> = (0..n)
.map(|m| ((2.0 * m as f64 + a * p) / (m as f64 + a)) as f32)
.collect();
assert_eq!(got, want);
}
#[test]
fn inference_uses_all_training_rows_and_prior_for_unseen() {
let nan = f32::NAN;
let cats = [0.0, 1.0, 0.0, nan, 1.0, 0.0];
let labels = [1.0, 4.0, 2.0, 9.0, 6.0, 0.5];
let enc = OrderedTargetEncoder::builder()
.prior_weight(2.0)
.build()
.unwrap();
let (train, fitted) = enc.fit_transform(&matrix(&cats, &labels), &[0]).unwrap();
let prior = labels.iter().map(|&y| f64::from(y)).sum::<f64>() / 6.0;
assert_eq!(fitted.prior(), prior as f32);
assert_eq!(train.get(3, 0), None);
let test = matrix(&[1.0, 0.0, 5.0, nan], &[0.0; 4]);
let out = fitted.transform(&test).unwrap();
let c0 = ((1.0 + 2.0 + 0.5) + 2.0 * prior) / (3.0 + 2.0);
let c1 = ((4.0 + 6.0) + 2.0 * prior) / (2.0 + 2.0);
assert_eq!(
column(&out, 0),
vec![Some(c1 as f32), Some(c0 as f32), Some(prior as f32), None]
);
assert_eq!(column(&out, 1), column(&test, 1));
assert_eq!(out.feature_types(), &[FeatureType::Numerical; 2]);
assert_eq!(fitted.encode(0, 5), Some(prior as f32));
assert_eq!(fitted.encode(1, 0), None);
}
#[test]
fn explicit_labels_encode_a_multi_target_matrix() {
let (cats, labels) = crafted();
let second: Vec<f32> = labels.iter().map(|&y| 3.0 - y).collect();
let pairs: Vec<f32> = labels
.iter()
.zip(&second)
.flat_map(|(&a, &b)| [a, b])
.collect();
let multi = matrix(&cats, &labels).with_label_matrix(&pairs, 2).unwrap();
let enc = encoder(2, 9);
assert!(matches!(
enc.fit_transform(&multi, &[0]),
Err(HessboostError::InvalidData {
input: "labels",
..
})
));
let (encoded, fitted) = enc
.fit_transform_with_labels(&multi, &[0], &second)
.unwrap();
let (single, single_fitted) = enc.fit_transform(&matrix(&cats, &second), &[0]).unwrap();
assert_eq!(column(&encoded, 0), column(&single, 0));
assert_eq!(fitted, single_fitted);
assert_eq!(encoded.n_targets(), 2);
assert_eq!(encoded.labels().unwrap(), pairs.as_slice());
assert!(matches!(
enc.fit_transform_with_labels(&multi, &[0], &second[1..]),
Err(HessboostError::DimensionMismatch { .. })
));
let mut bad = second.clone();
bad[3] = f32::NAN;
assert!(matches!(
enc.fit_transform_with_labels(&multi, &[0], &bad),
Err(HessboostError::InvalidData {
input: "labels",
..
})
));
}
#[test]
fn seed_determines_the_training_encoding() {
let (cats, labels) = crafted();
let data = matrix(&cats, &labels);
let run = |seed| column(&encoder(2, seed).fit_transform(&data, &[0]).unwrap().0, 0);
assert_eq!(run(5), run(5));
assert_ne!(run(5), run(6));
}
#[test]
fn many_permutations_converge_to_leave_one_out() {
let n = 20;
let cats = vec![0.0; n];
let labels: Vec<f32> = (0..n).map(|i| (i % 2) as f32).collect();
let separated = |permutations| {
let (out, _) = encoder(permutations, 1)
.fit_transform(&matrix(&cats, &labels), &[0])
.unwrap();
let enc = column(&out, 0);
let max_pos = (0..n)
.filter(|&i| labels[i] == 1.0)
.map(|i| enc[i].unwrap());
let min_neg = (0..n)
.filter(|&i| labels[i] == 0.0)
.map(|i| enc[i].unwrap());
max_pos.fold(f32::MIN, f32::max) < min_neg.fold(f32::MAX, f32::min)
};
assert!(separated(20_000));
assert!(!separated(1));
}
#[test]
fn csr_and_dense_inputs_encode_identically() {
let nan = f32::NAN;
let cats = [2.0, nan, 2.0, 7.0, nan, 7.0, 2.0];
let labels = [1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 0.0];
let dense = matrix(&cats, &labels);
let mut indptr = vec![0];
let (mut indices, mut values) = (Vec::new(), Vec::new());
for (i, &c) in cats.iter().enumerate() {
if !c.is_nan() {
indices.push(0);
values.push(c);
}
indices.push(1);
values.push(i as f32);
indptr.push(values.len());
}
let csr = DMatrix::from_csr(indptr, indices, values, 2)
.unwrap()
.with_labels(&labels)
.unwrap()
.with_feature_types(&CAT_NUM)
.unwrap();
let enc = encoder(3, 9);
let (d, fd) = enc.fit_transform(&dense, &[0]).unwrap();
let (s, fs) = enc.fit_transform(&csr, &[0]).unwrap();
assert_eq!(fd, fs);
for col in 0..2 {
assert_eq!(column(&d, col), column(&s, col));
}
assert_eq!(column(&s, 0)[1], None);
assert!(s.csr_parts().is_some());
}
#[test]
fn non_nan_missing_sentinel_is_preserved_and_cannot_collide() {
let x = [3.0, 0.0, 3.0, 5.0, 0.0, 1.0];
let data = DMatrix::from_dense_with_missing(&x, 3, 2, 0.0)
.unwrap()
.with_labels(&[0.0, 0.0, 1.0])
.unwrap()
.with_feature_types(&CAT_NUM)
.unwrap();
let enc = OrderedTargetEncoder::builder().prior(0.0).build().unwrap();
let (out, fitted) = enc.fit_transform(&data, &[0]).unwrap();
assert_eq!(column(&out, 1), vec![None, Some(5.0), Some(1.0)]);
assert_eq!(out.get(0, 0), Some(0.0));
assert_eq!(fitted.transform(&data).unwrap().get(0, 0), Some(0.0));
}
#[test]
fn serde_round_trip_preserves_transform_and_rejects_corruption() {
let (cats, labels) = crafted();
let data = matrix(&cats, &labels);
let (_, fitted) = encoder(1, 0).fit_transform(&data, &[0]).unwrap();
let json = serde_json::to_string(&fitted).unwrap();
let back: FittedTargetEncoder = serde_json::from_str(&json).unwrap();
assert_eq!(back, fitted);
assert_eq!(
column(&back.transform(&data).unwrap(), 0),
column(&fitted.transform(&data).unwrap(), 0)
);
let value: serde_json::Value = serde_json::from_str(&json).unwrap();
let rejects = |corrupt: fn(&mut serde_json::Value)| {
let mut v = value.clone();
corrupt(&mut v["columns"][0]);
serde_json::from_value::<FittedTargetEncoder>(v).is_err()
};
assert!(rejects(|c| c["categories"][0] = 2.into()), "unsorted");
assert!(
rejects(|c| c["values"] = serde_json::json!([1.0])),
"lengths"
);
assert!(rejects(|c| c["column"] = 2.into()), "column out of range");
}
#[test]
fn deserialization_does_not_allocate_from_the_declared_width() {
let doc = |columns: &str| {
format!(
r#"{{"n_cols": {}, "prior": 0.5, "columns": [{columns}]}}"#,
usize::MAX
)
};
let column = r#"{"column": 7, "categories": [0, 1], "values": [0.25, 0.75]}"#;
let wide: FittedTargetEncoder = serde_json::from_str(&doc(column)).unwrap();
assert_eq!(wide.n_cols, usize::MAX);
assert!(
serde_json::from_str::<FittedTargetEncoder>(&doc(&format!("{column}, {column}")))
.is_err()
);
}
#[test]
fn extreme_prior_weights_keep_the_smoothed_mean_finite() {
let data = matrix(&[0.0], &[2.0]);
for (a, prior, first, fitted_value) in [
(1e308, None, 2.0, 2.0),
(f64::from_bits(1), Some(0.25), 0.25, 2.0),
] {
let mut builder = OrderedTargetEncoder::builder().prior_weight(a);
if let Some(p) = prior {
builder = builder.prior(p);
}
let (out, fitted) = builder.build().unwrap().fit_transform(&data, &[0]).unwrap();
assert_eq!(out.get(0, 0), Some(first), "a = {a:e}");
assert_eq!(fitted.encode(0, 0), Some(fitted_value), "a = {a:e}");
let json = serde_json::to_string(&fitted).unwrap();
assert_eq!(
serde_json::from_str::<FittedTargetEncoder>(&json).unwrap(),
fitted
);
}
assert!(
OrderedTargetEncoder::builder()
.prior(1e300)
.build()
.is_err()
);
}
#[test]
fn invalid_inputs_are_rejected() {
let (cats, labels) = crafted();
let data = matrix(&cats, &labels);
let enc = encoder(1, 0);
let unlabeled = DMatrix::from_dense(&[0.0, 1.0], 1, 2)
.unwrap()
.with_feature_types(&CAT_NUM)
.unwrap();
assert!(enc.fit_transform(&unlabeled, &[0]).is_err());
assert!(enc.fit_transform(&data, &[]).is_err());
assert!(enc.fit_transform(&data, &[1]).is_err(), "numeric column");
assert!(enc.fit_transform(&data, &[0, 0]).is_err());
assert!(enc.fit_transform(&data, &[2]).is_err());
let binary = OrderedTargetEncoder::builder()
.target(TargetKind::Binary)
.build()
.unwrap();
let multiclass = matrix(&[0.0, 1.0, 1.0], &[0.0, 2.0, 1.0]);
assert!(binary.fit_transform(&multiclass, &[0]).is_err());
assert!(
OrderedTargetEncoder::builder()
.prior_weight(0.0)
.build()
.is_err()
);
assert!(
OrderedTargetEncoder::builder()
.permutations(0)
.build()
.is_err()
);
assert!(
OrderedTargetEncoder::builder()
.prior(f64::NAN)
.build()
.is_err()
);
let (_, fitted) = enc.fit_transform(&data, &[0]).unwrap();
let wide = DMatrix::from_dense(&[0.0; 3], 1, 3).unwrap();
assert!(fitted.transform(&wide).is_err());
let numeric = DMatrix::from_dense(&[0.0, 1.0], 1, 2).unwrap();
assert!(fitted.transform(&numeric).is_err());
}
#[test]
fn priors_stay_within_the_f32_range() {
let above = f64::from_bits(0x47ef_ffff_efff_ffff);
assert!(above > f64::from(f32::MAX) && (above as f32) == f32::MAX);
let with_prior = |prior: f64| {
OrderedTargetEncoder::builder()
.prior(prior)
.permutations(105)
.build()
};
assert!(matches!(
with_prior(above),
Err(HessboostError::InvalidParameter { .. })
));
let data = matrix(&[0.0, 1.0, 1.0, 2.0], &[1.0, 0.0, 1.0, 0.0]);
for prior in [f64::from(f32::MAX), -f64::from(f32::MAX)] {
let (encoded, fitted) = with_prior(prior)
.unwrap()
.fit_transform(&data, &[0])
.unwrap();
assert!(column(&encoded, 0).iter().all(|v| v.unwrap().is_finite()));
assert!(fitted.encode(0, 7).unwrap().is_finite());
}
}
}