mod common;
use common::bits::bits;
use std::num::NonZeroUsize;
use std::sync::mpsc::{Sender, channel};
use hessboost::config::{
BoosterKind, Dart, Langevin, LinearTree, ModelShrink, ModelShrinkMode, Monotone, MultiStrategy,
QuantizedGrad, TrainingParamsBuilder,
};
use hessboost::metric::Metric;
use hessboost::model::Predictions;
use hessboost::objective::distributional::{DistFamily, Distributional};
use hessboost::objective::{Multiclass, Objective, RegLoss};
use hessboost::prelude::*;
use serde_json::json;
fn regression(n: usize) -> DMatrix {
let mut noise = common::lcg(17);
let mut x = Vec::with_capacity(n * 4);
let mut y = Vec::with_capacity(n);
for i in 0..n {
let row = common::four_features(i);
x.extend_from_slice(&row);
y.push(3.0 * row[0] - 2.0 * row[1] * row[2] + noise() - 0.5);
}
common::labeled_dense(&x, 4, &y)
}
fn classification(n: usize, classes: usize) -> DMatrix {
let mut noise = common::lcg(5);
let mut x = Vec::with_capacity(n * 4);
let mut y = Vec::with_capacity(n);
for i in 0..n {
let row = common::four_features(i);
x.extend_from_slice(&row);
let score = row[0] + 0.5 * row[1] + 0.3 * noise();
y.push(((score * classes as f32 / 1.8) as usize).min(classes - 1) as f32);
}
common::labeled_dense(&x, 4, &y)
}
fn shrink(rate: f64, mode: ModelShrinkMode) -> ModelShrink {
ModelShrink::new(rate, mode).unwrap()
}
#[test]
fn truncations_are_the_shorter_runs_bit_for_bit() {
let base = || TrainingParams::builder().max_depth(3).eta(0.2).seed(9);
let cases: Vec<(&str, TrainingParams, DMatrix)> = vec![
(
"hist posterior sampling",
base()
.tree_method(TreeMethod::Hist)
.posterior_sampling(true)
.subsample(0.8)
.build()
.unwrap(),
regression(300),
),
(
"exact decreasing shrinkage with Langevin",
base()
.tree_method(TreeMethod::Exact)
.objective(Objective::BinaryLogistic(RegLoss::default()))
.langevin(
Langevin::builder()
.diffusion_temperature(50.0)
.build()
.unwrap(),
)
.model_shrink(shrink(0.3, ModelShrinkMode::Decreasing))
.build()
.unwrap(),
classification(300, 2),
),
(
"approx multiclass posterior sampling",
base()
.tree_method(TreeMethod::Approx)
.objective(Objective::Softprob(Multiclass::new(3).unwrap()))
.posterior_sampling(true)
.build()
.unwrap(),
classification(300, 3),
),
(
"vector-leaf dist:normal posterior sampling",
base()
.tree_method(TreeMethod::Hist)
.objective(Objective::Dist(Distributional::new(DistFamily::Normal)))
.multi_strategy(MultiStrategy::MultiOutputTree)
.posterior_sampling(true)
.build()
.unwrap(),
regression(300),
),
(
"shrinkage alone, boosted forest",
base()
.tree_method(TreeMethod::Hist)
.model_shrink(shrink(0.5, ModelShrinkMode::Constant))
.num_parallel_tree(2)
.subsample(0.7)
.build()
.unwrap(),
regression(300),
),
];
let rounds = 24;
for (name, params, data) in &cases {
let long = train(params, data, rounds).unwrap();
let members = long.predict_virtual_ensembles(data, 4).unwrap();
assert_eq!(members.iterations(), &[15, 18, 21, 24], "{name}");
for k in [1, 7, 15, 18, 23, rounds] {
let short = train(params, data, k).unwrap();
let expected = bits(short.predict_margin(data, Iterations::Best).unwrap());
assert_eq!(
bits(long.predict_margin(data, ..k).unwrap()),
expected,
"{name}: iterations ..{k}"
);
let sliced = long.slice(..k, 1).unwrap();
assert_eq!(
sliced.encode(ModelFormat::Json).unwrap(),
short.encode(ModelFormat::Json).unwrap(),
"{name}: slice ..{k}"
);
if let Some(m) = members.iterations().iter().position(|&it| it == k) {
assert_eq!(
bits(members.member_margins(m).unwrap()),
expected,
"{name}: member {m}"
);
}
}
}
}
#[test]
fn virtual_ensemble_members_are_the_prefix_predictions() {
let base = || {
TrainingParams::builder()
.max_depth(3)
.eta(0.3)
.seed(3)
.tree_method(TreeMethod::Hist)
};
let softmax = || Objective::Softmax(Multiclass::new(3).unwrap());
let cases: Vec<(&str, TrainingParamsBuilder, DMatrix)> = vec![
("gbtree", base(), regression(300)),
(
"constant shrinkage",
base().model_shrink(shrink(0.4, ModelShrinkMode::Constant)),
regression(300),
),
(
"decreasing shrinkage, linear leaves",
base()
.model_shrink(shrink(0.3, ModelShrinkMode::Decreasing))
.linear_tree(LinearTree::default()),
regression(300),
),
(
"linear leaves",
base().linear_tree(LinearTree::default()),
regression(300),
),
(
"dart",
base().booster(BoosterKind::Dart(Dart::default())),
regression(300),
),
(
"boosted forest",
base().num_parallel_tree(2).subsample(0.7),
regression(300),
),
(
"vector-leaf softmax",
base()
.objective(softmax())
.multi_strategy(MultiStrategy::MultiOutputTree),
classification(300, 3),
),
(
"softmax, decreasing shrinkage",
base()
.objective(softmax())
.model_shrink(shrink(0.3, ModelShrinkMode::Decreasing)),
classification(300, 3),
),
("wide CSR", base(), wide_csr_regression(300)),
];
for (name, params, train_data) in cases {
let model = train(¶ms.build().unwrap(), &train_data, 20).unwrap();
let k = model.n_outputs();
let few = train_data.select_rows(&[0, 1, 2, 3, 4]).unwrap();
let offsets: Vec<f32> = (0..300 * k).map(|i| (i % 13) as f32 * 0.25).collect();
let offset = train_data.clone().with_base_margin(&offsets).unwrap();
for (probe, data) in [
("all rows", &train_data),
("5 rows", &few),
("base_margin", &offset),
] {
let members = model.predict_virtual_ensembles(data, 4).unwrap();
for (m, &end) in members.iterations().iter().enumerate() {
assert_eq!(
bits(members.member_margins(m).unwrap()),
bits(model.predict_margin(data, ..end).unwrap()),
"{name}, {probe}: member {m} margins"
);
assert_eq!(
bits(members.member_predictions(m).unwrap()),
bits(model.predict(data, ..end).unwrap()),
"{name}, {probe}: member {m} predictions"
);
}
}
}
}
fn wide_csr_regression(n: usize) -> DMatrix {
let dense = regression(n);
let labels = dense.labels().unwrap().to_vec();
let (mut indptr, mut indices, mut values) = (vec![0], Vec::new(), Vec::new());
for i in 0..n {
for (f, v) in common::four_features(i).into_iter().enumerate() {
if !v.is_nan() {
indices.push(f as u32);
values.push(v);
}
}
indptr.push(indices.len());
}
DMatrix::from_csr(indptr, indices, values, 4 + 4096)
.unwrap()
.with_labels(&labels)
.unwrap()
}
#[test]
fn early_stopping_keeps_the_best_iteration_model() {
let data = regression(400);
let valid = regression(120);
let params = TrainingParams::builder()
.max_depth(4)
.eta(0.5)
.posterior_sampling(true)
.build()
.unwrap();
let result = Trainer::new(¶ms, &data, 200)
.eval(&valid, "valid")
.early_stopping_rounds(NonZeroUsize::new(3).unwrap())
.train()
.unwrap();
let model = result.model;
let best = model.best_iteration().unwrap();
assert!(best + 1 < result.history.len(), "training must stop early");
assert_eq!(model.num_boost_rounds(), best + 1);
let short = train(¶ms, &data, best + 1).unwrap();
assert_eq!(
bits(model.predict_margin(&valid, Iterations::Best).unwrap()),
bits(short.predict_margin(&valid, Iterations::Best).unwrap())
);
}
struct Recorder(Sender<Vec<f32>>);
impl Metric for Recorder {
fn name(&self) -> &'static str {
"recorded"
}
fn eval(&self, preds: &[f32], _labels: &[f32], _weights: Option<&[f32]>) -> f64 {
self.0.send(preds.to_vec()).unwrap();
0.0
}
}
#[test]
fn predictions_are_the_training_margins() {
let one = DMatrix::from_dense(&[0.0], 1, 1)
.unwrap()
.with_labels(&[0.0])
.unwrap();
let params = TrainingParams::builder()
.base_score(100_663_296.0)
.eta(1.0)
.lambda(0.0)
.model_shrink(shrink(0.7, ModelShrinkMode::Constant))
.build()
.unwrap();
let result = Trainer::new(¶ms, &one, 2)
.eval(&one, "train")
.train()
.unwrap();
assert_eq!(result.history.last().unwrap().values()[0], 0.0);
assert_eq!(
bits(result.model.predict_margin(&one, Iterations::Best).unwrap()),
bits([0.0])
);
let base = || TrainingParams::builder().max_depth(3).eta(0.3).seed(4);
let targets: Vec<f32> = (0..300)
.flat_map(|i| [(i % 7) as f32 * 1e3, (i % 5) as f32 - 2.0])
.collect();
let x: Vec<f32> = (0..300).flat_map(common::four_features).collect();
let matrix = DMatrix::from_dense(&x, 300, 4)
.unwrap()
.with_label_matrix(&targets, 2)
.unwrap();
let cases: Vec<(&str, TrainingParams, DMatrix)> = vec![
(
"hist posterior sampling",
base()
.base_score(1234.5)
.posterior_sampling(true)
.build()
.unwrap(),
regression(300),
),
(
"exact decreasing shrinkage",
base()
.tree_method(TreeMethod::Exact)
.model_shrink(shrink(0.4, ModelShrinkMode::Decreasing))
.build()
.unwrap(),
regression(300),
),
(
"shrinkage alone, boosted forest",
base()
.model_shrink(shrink(0.5, ModelShrinkMode::Constant))
.num_parallel_tree(2)
.subsample(0.7)
.build()
.unwrap(),
regression(300),
),
(
"shrinkage alone, linear leaves",
base()
.model_shrink(shrink(0.5, ModelShrinkMode::Constant))
.linear_tree(LinearTree::default())
.build()
.unwrap(),
regression(300),
),
(
"vector leaves over a label matrix",
base()
.multi_strategy(MultiStrategy::MultiOutputTree)
.posterior_sampling(true)
.build()
.unwrap(),
matrix.clone(),
),
(
"scalar trees over a label matrix",
base().posterior_sampling(true).build().unwrap(),
matrix,
),
];
let rounds = 12;
for (name, params, data) in cases {
let (tx, rx) = channel();
let model = Trainer::new(¶ms, &data, rounds)
.eval(&data, "train")
.custom_metric(Box::new(Recorder(tx)))
.train()
.unwrap()
.model;
let seen: Vec<Vec<f32>> = rx.try_iter().collect();
assert_eq!(seen.len(), rounds, "{name}");
for (k, margins) in seen.iter().enumerate() {
assert_eq!(
bits(model.predict_margin(&data, ..=k).unwrap()),
bits(margins),
"{name}: after round {k}"
);
}
assert_eq!(
bits(model.predict_margin(&data, Iterations::Best).unwrap()),
bits(&seen[rounds - 1]),
"{name}"
);
}
}
#[test]
fn langevin_training_is_thread_count_independent() {
let data = regression(20_000);
let params = |seed| {
TrainingParams::builder()
.objective(Objective::SquaredError(RegLoss::default()))
.tree_method(TreeMethod::Hist)
.max_depth(4)
.posterior_sampling(true)
.seed(seed)
.build()
.unwrap()
};
let run = |threads, seed| {
common::with_threads(threads, || {
train(¶ms(seed), &data, 6)
.unwrap()
.encode(ModelFormat::Binary)
.unwrap()
})
};
let serial = run(1, 0);
assert_eq!(serial, run(4, 0));
assert_ne!(serial, run(1, 1));
let plain = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(4)
.build()
.unwrap();
assert_ne!(
serial,
train(&plain, &data, 6)
.unwrap()
.encode(ModelFormat::Binary)
.unwrap()
);
}
#[test]
fn renewed_leaves_keep_min_child_weight() {
let x = [0.0, 1.0, 2.0, 3.0];
let scalar = common::labeled_dense(&x, 1, &[5.0, 6.0, 7.0, 8.0]);
let matrix = DMatrix::from_dense(&x, 4, 1)
.unwrap()
.with_label_matrix(&[5.0, -1.0, 6.0, -2.0, 7.0, -3.0, 8.0, -4.0], 2)
.unwrap();
let params = |strategy, langevin: bool| {
let builder = TrainingParams::builder()
.base_score(0.5)
.min_child_weight(10.0)
.multi_strategy(strategy)
.model_shrink(shrink(0.1, ModelShrinkMode::Constant));
let builder = if langevin {
builder.langevin(Langevin::default())
} else {
builder
};
builder.build().unwrap()
};
for (strategy, data) in [
(MultiStrategy::OneOutputPerTree, &scalar),
(MultiStrategy::OneOutputPerTree, &matrix),
(MultiStrategy::MultiOutputTree, &matrix),
] {
let noisy = train(¶ms(strategy, true), data, 3).unwrap();
let plain = train(¶ms(MultiStrategy::OneOutputPerTree, false), data, 3).unwrap();
assert_eq!(
bits(noisy.predict_margin(data, Iterations::Best).unwrap()),
bits(plain.predict_margin(data, Iterations::Best).unwrap()),
"{strategy:?}"
);
}
}
#[test]
fn shrunk_models_round_trip() {
let data = regression(200);
let params = TrainingParams::builder()
.max_depth(3)
.posterior_sampling(true)
.build()
.unwrap();
let model = train(¶ms, &data, 12).unwrap();
let margin = model.predict_margin(&data, Iterations::Best).unwrap();
let margins = bits(&margin);
let prefix = bits(model.predict_margin(&data, ..5).unwrap());
for restored in [
BoostedModel::decode(
model.encode(ModelFormat::Binary).unwrap(),
ModelFormat::Binary,
)
.unwrap(),
BoostedModel::decode(model.encode(ModelFormat::Json).unwrap(), ModelFormat::Json).unwrap(),
] {
assert_eq!(
bits(restored.predict_margin(&data, Iterations::Best).unwrap()),
margins
);
assert_eq!(bits(restored.predict_margin(&data, ..5).unwrap()), prefix);
}
let xgboost = BoostedModel::decode(
model.encode(ModelFormat::XgboostJson).unwrap(),
ModelFormat::XgboostJson,
)
.unwrap();
for (&x, &m) in xgboost
.predict_margin(&data, Iterations::Best)
.unwrap()
.as_slice()
.iter()
.zip(margin.as_slice())
{
assert!((x - m).abs() <= 1e-5 * m.abs().max(1.0), "{x} vs {m}");
}
let compact = model.to_compact().unwrap();
assert_eq!(bits(compact.predict_margin(&data).unwrap()), margins);
let compact = hessboost::model::compact::CompactModel::decode(compact.encode()).unwrap();
assert_eq!(bits(compact.predict_margin(&data).unwrap()), margins);
let offset = regression(200).with_base_margin(&[0.5; 200]).unwrap();
assert_eq!(
bits(compact.predict_margin(&offset).unwrap()),
bits(model.predict_margin(&offset, Iterations::Best).unwrap())
);
let contribs = model.predict_contribs(&data, Iterations::Best).unwrap();
for (row, &m) in margin.as_slice().iter().enumerate() {
let sum: f32 = contribs.get(row, 0).unwrap().iter().sum();
assert!((sum - m).abs() < 1e-4, "{sum} vs {m}");
}
let mut doc: serde_json::Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
doc["shrinkage"]["factors"][3] = 0.5.into();
let err = BoostedModel::decode(doc.to_string(), ModelFormat::Json).unwrap_err();
assert!(
matches!(&err, HessboostError::ModelFormat(msg) if msg.contains("shrinkage record")),
"{err:?}"
);
}
#[test]
fn shrunk_models_refuse_non_prefix_selections() {
let data = regression(100);
let params = TrainingParams::builder()
.max_depth(2)
.model_shrink(shrink(0.1, ModelShrinkMode::Constant))
.build()
.unwrap();
let model = train(¶ms, &data, 10).unwrap();
assert_eq!(
common::incompatible_model(model.predict_margin(&data, 2..5)),
"iterations"
);
assert_eq!(common::incompatible_model(model.slice(2..5, 1)), "slice");
assert_eq!(common::incompatible_model(model.slice(..6, 2)), "slice");
assert_eq!(
common::incompatible_model(model.predict_contribs(&data, ..4)),
"iterations"
);
assert!(model.predict_contribs(&data, ..).is_ok());
assert!(model.predict_leaf(&data, ..4).is_ok());
}
#[test]
fn langevin_refuses_quantized_leaf_renewal() {
let renewed = QuantizedGrad::builder().renew_leaf(true).build().unwrap();
for builder in [
TrainingParams::builder().langevin(Langevin::default()),
TrainingParams::builder().posterior_sampling(true),
] {
let quantized = builder.clone().quantized(QuantizedGrad::default());
assert!(quantized.build().is_ok());
assert_eq!(
common::invalid_param(builder.quantized(renewed).build()),
"langevin"
);
}
let flat = [
("posterior_sampling", json!(true)),
("use_quantized_grad", json!(true)),
("quant_train_renew_leaf", json!(true)),
];
assert_eq!(
common::invalid_param(TrainingParams::from_xgboost(flat)),
"langevin"
);
}
#[test]
fn langevin_noise_scale_must_be_representable() {
let params = |eta: f64, temperature: f64| {
let langevin = Langevin::builder()
.diffusion_temperature(temperature)
.build()
.unwrap();
TrainingParams::builder()
.eta(eta)
.langevin(langevin)
.build()
};
for (eta, temperature) in [(1e-10, 1e-300), (1e30, 1e300), (1.0, 1e-78), (1.0, 1e92)] {
assert_eq!(
common::invalid_param(params(eta, temperature)),
"diffusion_temperature",
"eta {eta}, temperature {temperature}"
);
}
assert!(params(1.0, 1e-70).is_ok() && params(1.0, 1e70).is_ok());
assert!(params(1.0, 1e78).is_ok());
let tiny = TrainingParams::builder()
.eta(1e-44)
.posterior_sampling(true)
.build()
.unwrap();
assert!(train(&tiny, ®ression(2), 1).is_ok());
}
#[test]
fn unsupported_combinations_are_refused() {
let refused = |builder: TrainingParamsBuilder| common::invalid_param(builder.build());
let base = TrainingParams::builder;
let tempered = || {
Langevin::builder()
.diffusion_temperature(10.0)
.build()
.unwrap()
};
let constant = |rate| shrink(rate, ModelShrinkMode::Constant);
assert_eq!(
refused(base().posterior_sampling(true).langevin(tempered())),
"diffusion_temperature"
);
assert_eq!(
refused(base().posterior_sampling(true).model_shrink(constant(0.01))),
"model_shrink_rate"
);
assert!(
base()
.posterior_sampling(true)
.langevin(Langevin::default())
.build()
.is_ok()
);
assert_eq!(
refused(base().eta(0.5).model_shrink(constant(2.0))),
"model_shrink_rate"
);
let dart = BoosterKind::Dart(Dart::default());
assert_eq!(
refused(base().langevin(Langevin::default()).booster(dart)),
"langevin"
);
assert_eq!(
refused(base().model_shrink(constant(0.1)).booster(dart)),
"model_shrink_rate"
);
assert_eq!(
refused(
base()
.model_shrink(constant(0.1))
.booster(BoosterKind::GbLinear)
),
"model_shrink_rate"
);
for builder in [
base().num_parallel_tree(2),
base().monotone_constraints(vec![Monotone::Increasing]),
base().path_smooth(1.0),
base().linear_tree(LinearTree::default()),
] {
assert_eq!(refused(builder.langevin(Langevin::default())), "langevin");
}
let flat = |pairs: &[(&str, serde_json::Value)]| {
common::invalid_param(TrainingParams::from_xgboost(pairs.iter().cloned()))
};
assert_eq!(
flat(&[("diffusion_temperature", json!(10.0))]),
"diffusion_temperature"
);
assert_eq!(
flat(&[("model_shrink_mode", json!("decreasing"))]),
"model_shrink_mode"
);
assert_eq!(
flat(&[
("posterior_sampling", json!(true)),
("langevin", json!(false))
]),
"langevin"
);
let round_trip = base()
.langevin(tempered())
.model_shrink(shrink(0.2, ModelShrinkMode::Decreasing))
.build()
.unwrap();
assert_eq!(
TrainingParams::from_xgboost(round_trip.to_xgboost().unwrap()).unwrap(),
round_trip
);
let data = regression(50);
let params = base().posterior_sampling(true).build().unwrap();
let with_margin = regression(50).with_base_margin(&[0.5; 50]).unwrap();
assert_eq!(
common::invalid_data(train(¶ms, &with_margin, 2)),
("base_margin", Some("dtrain".into()))
);
let shrunk = train(¶ms, &data, 4).unwrap();
assert_eq!(
common::incompatible_model(Trainer::new(¶ms, &data, 2).init_model(&shrunk).train()),
"init_model"
);
let plain_params = base().build().unwrap();
assert_eq!(
common::incompatible_model(
Trainer::new(&plain_params, &data, 2)
.init_model(&shrunk)
.train()
),
"init_model"
);
let steep = base().posterior_sampling(true).eta(3.0).build().unwrap();
assert_eq!(
common::invalid_param(train(&steep, ®ression(1), 1)),
"posterior_sampling"
);
assert!(train(&steep, ®ression(2), 1).is_ok());
}
#[test]
fn no_model_shrinkage_is_none() {
assert_eq!(
common::invalid_param(ModelShrink::new(0.0, ModelShrinkMode::Constant)),
"model_shrink_rate"
);
let data = regression(100);
let typed = TrainingParams::builder()
.langevin(Langevin::default())
.build()
.unwrap();
let model = train(&typed, &data, 3).unwrap();
assert!(
Trainer::new(&typed, &data, 2)
.init_model(&model)
.train()
.is_ok()
);
let flat =
|pairs: &[(&str, serde_json::Value)]| TrainingParams::from_xgboost(pairs.iter().cloned());
let catboost = flat(&[("langevin", json!(true))]).unwrap();
assert_eq!(
catboost.model_shrink,
Some(shrink(0.001, ModelShrinkMode::Constant))
);
let off = flat(&[("langevin", json!(true)), ("model_shrink_rate", json!(0.0))]).unwrap();
assert_eq!(off, typed);
for params in [&typed, &catboost] {
assert_eq!(
&TrainingParams::from_xgboost(params.to_xgboost().unwrap()).unwrap(),
params
);
}
assert_eq!(
common::invalid_param(flat(&[
("model_shrink_rate", json!(0.0)),
("model_shrink_mode", json!("decreasing"))
])),
"model_shrink_mode"
);
assert_eq!(
common::invalid_param(flat(&[
("posterior_sampling", json!(true)),
("model_shrink_rate", json!(0.0))
])),
"model_shrink_rate"
);
}
#[test]
fn langevin_continuation_matches_the_uninterrupted_run() {
let data = regression(200);
let params = TrainingParams::builder()
.max_depth(3)
.langevin(Langevin::default())
.build()
.unwrap();
let first = train(¶ms, &data, 5).unwrap();
let continued = Trainer::new(¶ms, &data, 4)
.init_model(&first)
.train()
.unwrap()
.model;
let whole = train(¶ms, &data, 9).unwrap();
assert_eq!(
bits(continued.predict_margin(&data, Iterations::Best).unwrap()),
bits(whole.predict_margin(&data, Iterations::Best).unwrap())
);
}
#[test]
fn uncertainty_decomposes_per_objective() {
let binary = classification(300, 2);
let params = |objective: Objective| {
TrainingParams::builder()
.objective(objective)
.max_depth(3)
.posterior_sampling(true)
.build()
.unwrap()
};
let model = train(
¶ms(Objective::BinaryLogistic(RegLoss::default())),
&binary,
40,
)
.unwrap();
let u = model.predict_uncertainty(&binary, 10).unwrap();
let (data, total) = (u.data.unwrap(), u.total.unwrap());
let cells = |p: &Predictions<f64>| p.as_slice().to_vec();
for ((&k, &d), &t) in cells(&u.knowledge)
.iter()
.zip(&cells(&data))
.zip(&cells(&total))
{
assert!((k - (t - d)).abs() < 1e-15);
assert!(k > -1e-12 && d >= 0.0 && t <= std::f64::consts::LN_2 + 1e-12);
}
assert!(u.mean.as_slice().iter().all(|&p| (0.0..=1.0).contains(&p)));
let reg = regression(300);
let dist = train(
¶ms(Objective::Dist(Distributional::new(DistFamily::Normal))),
®,
40,
)
.unwrap();
let u = dist.predict_uncertainty(®, 5).unwrap();
let (data, total) = (u.data.unwrap(), u.total.unwrap());
for ((&k, &d), &t) in cells(&u.knowledge)
.iter()
.zip(&cells(&data))
.zip(&cells(&total))
{
assert!(k >= 0.0 && d > 0.0);
assert_eq!(t, k + d);
}
let squared = train(
¶ms(Objective::SquaredError(RegLoss::default())),
®,
40,
)
.unwrap();
let u = squared.predict_uncertainty(®, 5).unwrap();
assert!(u.data.is_none() && u.total.is_none());
assert!(u.knowledge.as_slice().iter().all(|&k| k >= 0.0));
assert_eq!(
common::incompatible_model(squared.predict_virtual_ensembles(®, 21)),
"virtual_ensembles_count"
);
assert_eq!(
common::invalid_param(squared.predict_virtual_ensembles(®, 0)),
"virtual_ensembles_count"
);
assert_eq!(
common::incompatible_model(squared.predict_virtual_ensembles(®, usize::MAX)),
"virtual_ensembles_count"
);
let members = squared.predict_virtual_ensembles(®, 5).unwrap();
assert!(members.member_margins(4).is_some());
assert!(members.member_margins(5).is_none());
assert!(members.member_predictions(usize::MAX).is_none());
assert!(members.member_margins(usize::MAX).is_none());
}
#[test]
fn uncertainty_parts_have_their_own_widths() {
let params = |objective: Objective| {
TrainingParams::builder()
.objective(objective)
.max_depth(2)
.posterior_sampling(true)
.build()
.unwrap()
};
let labels = |i: usize| [(i % 2) as f32, ((i / 2) % 2) as f32];
let x: Vec<f32> = (0..120).flat_map(common::four_features).collect();
let multi_label = DMatrix::from_dense(&x, 120, 4)
.unwrap()
.with_label_matrix(&(0..120).flat_map(labels).collect::<Vec<_>>(), 2)
.unwrap();
let multiclass = classification(120, 3);
let reg = regression(120);
let cases = [
(
Objective::Softprob(Multiclass::new(3).unwrap()),
&multiclass,
3,
1,
),
(
Objective::Softmax(Multiclass::new(3).unwrap()),
&multiclass,
3,
1,
),
(
Objective::BinaryLogistic(RegLoss::default()),
&multi_label,
2,
2,
),
(
Objective::SquaredError(RegLoss::default()),
&multi_label,
2,
2,
),
(
Objective::Dist(Distributional::new(DistFamily::Normal)),
®,
1,
1,
),
];
for (objective, data, mean_width, width) in cases {
let name = objective.name().to_string();
let model = train(¶ms(objective), data, 20).unwrap();
let u = model.predict_uncertainty(data, 5).unwrap();
assert_eq!(
(u.mean.n_rows(), u.mean.width()),
(120, mean_width),
"{name}"
);
let parts = [Some(&u.knowledge), u.data.as_ref(), u.total.as_ref()];
for part in parts.into_iter().flatten() {
assert_eq!((part.n_rows(), part.width()), (120, width), "{name}");
}
}
}