use hessboost::config::{BoosterKind, Dart, LinearTree, MultiStrategy};
use hessboost::data::FeatureType;
use hessboost::diffusion::DiffusionFormat;
use hessboost::diffusion::forest::{ColumnKind, ForestModel, ForestParams, NoiseLevels};
use hessboost::diffusion::{
DiffusionModel, DiffusionParams, Method, SampleOptions, ScoreConfig, Sde,
};
use hessboost::model::compact::CompactModel;
use hessboost::objective::distributional::{
DistFamily, DistGradient, DistSplitDirection, Distributional,
};
use hessboost::objective::{Aft, AftDistribution, Expectiles, Multiclass, Quantiles};
use hessboost::prelude::*;
use serde_json::Value;
use std::path::{Path, PathBuf};
mod common;
use common::bits::bits;
use common::{four_features, labeled_dense};
const COLS: usize = 4;
#[test]
fn unknown_and_corrupt_native_payloads_are_refused() {
let model = train(&base().build().unwrap(), &matrix(1), 3).unwrap();
let bytes = model.encode(ModelFormat::Binary).unwrap();
let container = zstd::stream::decode_all(bytes.as_slice()).unwrap();
assert_eq!(&container[..4], b"HBM\0");
assert_eq!(
bits(
BoostedModel::decode(&container, ModelFormat::Binary)
.unwrap()
.predict(&matrix(1), Iterations::Best)
.unwrap()
.as_slice()
),
bits(
model
.predict(&matrix(1), Iterations::Best)
.unwrap()
.as_slice()
)
);
let with_version = |version: u8| {
let mut bytes = container.clone();
bytes[4] = version;
bytes
};
for corrupt in [
with_version(0),
with_version(1),
with_version(2),
with_version(4),
with_version(255),
bytes[..bytes.len() - 3].to_vec(),
container[..container.len() - 3].to_vec(),
container[..4].to_vec(),
] {
let err = BoostedModel::decode(&corrupt, ModelFormat::Binary).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
}
}
fn empty_model_doc() -> Value {
let model = train(&base().build().unwrap(), &matrix(1), 0).unwrap();
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap()
}
fn load_doc(doc: &Value) -> hessboost::error::Result<BoostedModel> {
BoostedModel::decode(doc.to_string(), ModelFormat::Json)
}
fn trained_doc(params: &TrainingParams, data: &DMatrix) -> Value {
let model = train(params, data, 4).unwrap();
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap()
}
fn assert_refused(doc: &Value, what: &str) {
let err = load_doc(doc).expect_err(what);
assert!(
matches!(
err,
HessboostError::Json(_) | HessboostError::ModelFormat(_)
),
"{what}: {err}"
);
}
#[test]
fn incomplete_json_documents_are_refused() {
let doc = trained_doc(&base().build().unwrap(), &matrix(1));
for field in [
"objective",
"base_score",
"num_class",
"n_features",
"n_outputs",
"n_targets",
"trees",
"tree_weights",
"num_parallel_tree",
"linear",
] {
let mut doc = doc.clone();
doc.as_object_mut().unwrap().remove(field);
assert_refused(&doc, field);
}
for field in ["nodes", "categories", "linear"] {
let mut doc = doc.clone();
doc["trees"][0].as_object_mut().unwrap().remove(field);
assert_refused(&doc, field);
}
let vector = trained_doc(
&base()
.multi_strategy(MultiStrategy::MultiOutputTree)
.build()
.unwrap(),
&matrix(3),
);
assert!(load_doc(&vector).unwrap().has_vector_leaves());
for field in ["size_leaf_vector", "leaf_vectors"] {
let mut doc = vector.clone();
doc["trees"][0].as_object_mut().unwrap().remove(field);
assert_refused(&doc, field);
}
let mut doc = trained_doc(&base().build().unwrap(), &matrix(2));
doc["trees"][0]
.as_object_mut()
.unwrap()
.remove("size_leaf_vector");
assert_refused(&doc, "multi-target size_leaf_vector");
let linear = trained_doc(
&base().linear_tree(LinearTree::default()).build().unwrap(),
&matrix(1),
);
let tree = linear["trees"]
.as_array()
.unwrap()
.iter()
.position(|t| t["linear"].is_object())
.expect("a tree with linear leaves");
let parts: Vec<String> = linear["trees"][tree]["linear"]
.as_object()
.unwrap()
.keys()
.cloned()
.collect();
let mut doc = linear.clone();
doc["trees"][tree].as_object_mut().unwrap().remove("linear");
assert_refused(&doc, "linear leaves");
for part in &parts {
let mut doc = linear.clone();
doc["trees"][tree]["linear"]
.as_object_mut()
.unwrap()
.remove(part);
assert_refused(&doc, part);
}
}
fn default_objective_params(objective: &str) -> Value {
let mut defaults = empty_model_doc()["objective_params"].clone();
assert_eq!(defaults["max_delta_step"], 0.0);
assert!(defaults["distribution"].is_null());
if objective == "count:poisson" {
defaults["max_delta_step"] = 0.7.into();
}
defaults["distribution"] = serde_json::to_value(DistFamily::from_objective(objective)).unwrap();
defaults
}
fn strip_defaults(doc: &mut Value) -> usize {
let single_output = doc["n_outputs"] == 1;
let defaults = default_objective_params(doc["objective"].as_str().unwrap());
let params = doc["objective_params"].as_object_mut().unwrap();
let before = params.len();
params.retain(|key, value| defaults[key] != *value);
let removed = before - params.len();
if params.is_empty() {
doc.as_object_mut().unwrap().remove("objective_params");
}
for tree in doc["trees"].as_array_mut().unwrap() {
let tree = tree.as_object_mut().unwrap();
if tree["size_leaf_vector"] == 0 {
tree.remove("leaf_vectors");
if single_output {
tree.remove("size_leaf_vector");
}
}
}
removed
}
#[test]
fn json_documents_may_omit_defaults() {
let (x, _) = train_data(1);
let n = x.len() / COLS;
let counts: Vec<f32> = (0..n).map(|i| ((i * 7) % 9) as f32).collect();
let cases = [
("reg:squarederror", base().build().unwrap(), matrix(1)),
(
"count:poisson",
base().objective(Objective::Poisson).build().unwrap(),
labeled_dense(&x, COLS, &counts),
),
(
"reg:quantileerror",
base()
.objective(Objective::Quantile(
Quantiles::new([0.1, 0.5, 0.9]).unwrap(),
))
.build()
.unwrap(),
matrix(1),
),
(
"dist:normal",
base()
.objective(Objective::Dist(Distributional::new(DistFamily::Normal)))
.build()
.unwrap(),
matrix(1),
),
];
for (name, params, data) in cases {
let model = train(¶ms, &data, 4).unwrap();
let mut doc: Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
let removed = strip_defaults(&mut doc);
let quantiles = name == "reg:quantileerror";
assert_eq!(removed, if quantiles { 11 } else { 12 }, "{name}");
assert!(doc["trees"][0].get("leaf_vectors").is_none(), "{name}");
let restored = load_doc(&doc).unwrap();
assert_eq!(restored.objective(), model.objective(), "{name}");
assert_eq!(
restored.encode(ModelFormat::Binary).unwrap(),
model.encode(ModelFormat::Binary).unwrap(),
"{name}"
);
assert_eq!(
bits(
restored
.predict(&data, Iterations::Best)
.unwrap()
.as_slice()
),
bits(model.predict(&data, Iterations::Best).unwrap().as_slice()),
"{name}"
);
let rewritten: Value =
serde_json::from_slice(&restored.encode(ModelFormat::Json).unwrap()).unwrap();
assert_eq!(
rewritten,
serde_json::from_slice::<Value>(&model.encode(ModelFormat::Json).unwrap()).unwrap(),
"{name}"
);
if quantiles {
doc.as_object_mut().unwrap().remove("objective_params");
assert_refused(&doc, "quantile_alpha");
}
}
}
#[test]
fn overflowing_tree_layout_is_refused() {
let mut doc = empty_model_doc();
doc["n_outputs"] = 2.into();
doc["n_targets"] = 2.into();
doc["base_score"] = serde_json::json!([0.0, 0.0]);
assert!(load_doc(&doc).is_ok());
doc["num_parallel_tree"] = (1u64 << 63).into();
for best_iteration in [Value::Null, 0.into()] {
doc["best_iteration"] = best_iteration;
assert!(matches!(
load_doc(&doc),
Err(HessboostError::ModelFormat(_))
));
}
}
#[test]
fn objective_width_must_match_the_stored_outputs() {
for (objective, key) in [
("reg:expectileerror", "expectile_alpha"),
("reg:quantileerror", "quantile_alpha"),
] {
let mut doc = empty_model_doc();
doc["objective"] = objective.into();
doc["objective_params"][key] = serde_json::json!([0.2, 0.8]);
let err = load_doc(&doc).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
doc["n_outputs"] = 2.into();
doc["base_score"] = serde_json::json!([0.0, 0.0]);
let model = load_doc(&doc).unwrap();
assert_eq!(model.n_outputs(), 2);
}
}
fn saved_doc(name: &str) -> Value {
let path = saved_dir("0.2.0").join(format!("{name}.json"));
serde_json::from_str(&std::fs::read_to_string(path).unwrap()).unwrap()
}
#[test]
fn documents_other_readers_refuse_are_refused() {
let mut doc = empty_model_doc();
doc["objective"] = "binary:logistic".into();
assert!(load_doc(&doc).is_ok());
doc["objective_params"]["scale_pos_weight"] = 0.into();
assert!(matches!(
load_doc(&doc),
Err(HessboostError::ModelFormat(_))
));
let mut doc = saved_doc("categorical_splits");
assert!(load_doc(&doc).is_ok());
let node = &mut doc["trees"][0]["nodes"][0];
assert_eq!(node["is_categorical"], true);
node["cat_end"] = node["cat_begin"].clone();
assert!(matches!(
load_doc(&doc),
Err(HessboostError::ModelFormat(_))
));
let mut doc = saved_doc("categorical_splits");
let nodes = doc["trees"][0]["nodes"].as_array_mut().unwrap();
let leaf = nodes.iter_mut().find(|n| n["left"] == -1).unwrap();
leaf["is_categorical"] = true.into();
assert!(matches!(
load_doc(&doc),
Err(HessboostError::ModelFormat(_))
));
let mut doc = saved_doc("dist_normal");
assert!(load_doc(&doc).is_ok());
for family in [Value::Null, "gamma".into()] {
doc["objective_params"]["distribution"] = family;
assert!(matches!(
load_doc(&doc),
Err(HessboostError::ModelFormat(_))
));
}
}
#[test]
fn training_refuses_to_return_a_model_that_would_not_load() {
let (x, _) = train_data(1);
let n = x.len() / COLS;
let y: Vec<f32> = (0..n).map(|i| (i % 7) as f32).collect();
let data = DMatrix::from_dense(&x, n, COLS)
.unwrap()
.with_labels(&y)
.unwrap()
.with_weights(&vec![f32::MAX; n])
.unwrap();
let err = train(&base().build().unwrap(), &data, 2).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
}
fn train_data(k: usize) -> (Vec<f32>, Vec<f32>) {
const N: usize = 160;
let mut x = Vec::with_capacity(N * COLS);
let mut y = Vec::with_capacity(N * k);
for i in 0..N {
let row = four_features(i);
let [a, b, c, _] = row;
x.extend(row);
for t in 0..k {
y.push(2.0 * a - b + t as f32 * c + 0.1);
}
}
(x, y)
}
fn matrix(k: usize) -> DMatrix {
let (x, y) = train_data(k);
let n = x.len() / COLS;
let data = DMatrix::from_dense(&x, n, COLS).unwrap();
if k == 1 {
data.with_labels(&y).unwrap()
} else {
data.with_label_matrix(&y, k).unwrap()
}
}
fn base() -> hessboost::config::TrainingParamsBuilder {
TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(3)
.nthread(1)
.seed(11)
}
type Has = fn(&BoostedModel) -> bool;
fn has_dart_weights(model: &BoostedModel) -> bool {
let doc: Value = serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
doc["tree_weights"]
.as_array()
.unwrap()
.iter()
.any(|w| w.as_f64() != Some(1.0))
}
fn feature_models() -> Vec<(&'static str, BoostedModel, DMatrix, Has)> {
let (x, _) = train_data(1);
let n = x.len() / COLS;
let classes: Vec<f32> = (0..n).map(|i| (i % 3) as f32).collect();
let lower: Vec<f32> = (0..n).map(|i| 1.0 + (i % 5) as f32).collect();
let upper: Vec<f32> = lower
.iter()
.enumerate()
.map(|(i, &l)| if i % 4 == 0 { f32::INFINITY } else { l })
.collect();
let aft = DMatrix::from_dense(&x, n, COLS)
.unwrap()
.with_label_bounds(&lower, &upper)
.unwrap();
let softmax = labeled_dense(&x, COLS, &classes);
let counts: Vec<f32> = (0..n).map(|i| ((i * 7) % 9) as f32).collect();
let negbinomial = labeled_dense(&x, COLS, &counts);
let mut cat_x = x.clone();
let mut cat_y = Vec::with_capacity(n);
for i in 0..n {
cat_x[i * COLS] = (i % 5) as f32;
cat_y.push([0.0, 2.0, -1.0, 3.0, 1.0][i % 5] + cat_x[i * COLS + 1]);
}
let mut types = vec![FeatureType::Numerical; COLS];
types[0] = FeatureType::Categorical;
let categorical = labeled_dense(&cat_x, COLS, &cat_y)
.with_feature_types(&types)
.unwrap();
let cases: Vec<(&str, TrainingParams, DMatrix, Has)> = vec![
(
"dart",
base()
.booster(BoosterKind::Dart(
Dart::builder().rate_drop(0.5).build().unwrap(),
))
.build()
.unwrap(),
matrix(1),
has_dart_weights,
),
(
"gblinear",
base().booster(BoosterKind::GbLinear).build().unwrap(),
matrix(1),
|m| m.num_trees() == 0,
),
(
"categorical splits",
base().build().unwrap(),
categorical,
|m| {
m.trees()
.iter()
.any(|t| t.nodes().iter().any(|node| node.is_categorical))
},
),
(
"vector leaves",
base()
.multi_strategy(MultiStrategy::MultiOutputTree)
.build()
.unwrap(),
matrix(3),
|m| m.has_vector_leaves(),
),
(
"linear leaves",
base().linear_tree(LinearTree::default()).build().unwrap(),
matrix(1),
|m| m.trees().iter().any(|t| t.linear_leaves().is_some()),
),
(
"multiclass forest",
base()
.objective(Objective::Softprob(Multiclass::new(3).unwrap()))
.num_parallel_tree(3)
.subsample(0.7)
.colsample_bynode(0.6)
.build()
.unwrap(),
softmax,
|m| m.num_parallel_tree() == 3 && m.trees_per_iteration() == 9,
),
("multi-target", base().build().unwrap(), matrix(2), |m| {
m.n_targets() == 2 && m.n_outputs() == 2
}),
(
"quantiles",
base()
.objective(Objective::Quantile(
Quantiles::new([0.1, 0.5, 0.9]).unwrap(),
))
.build()
.unwrap(),
matrix(1),
|m| m.n_outputs() == 3,
),
(
"expectiles",
base()
.objective(Objective::Expectile(Expectiles::new([0.2, 0.8]).unwrap()))
.build()
.unwrap(),
matrix(1),
|m| m.n_outputs() == 2,
),
(
"aft",
base()
.objective(Objective::Aft(
Aft::new(AftDistribution::Logistic, 0.7).unwrap(),
))
.build()
.unwrap(),
aft,
|m| {
matches!(m.objective().built_in(), Some(Objective::Aft(a))
if a.distribution() == AftDistribution::Logistic)
},
),
(
"dist:normal",
base()
.objective(Objective::Dist(
Distributional::new(DistFamily::Normal).with_gradient(DistGradient::Hessian),
))
.build()
.unwrap(),
matrix(1),
|m| {
matches!(m.objective().built_in(), Some(Objective::Dist(d))
if d.family() == DistFamily::Normal && d.gradient() == DistGradient::Hessian)
&& m.n_outputs() == 2
&& !m.has_vector_leaves()
},
),
(
"dist:normal vector leaves",
base()
.objective(Objective::Dist(
Distributional::new(DistFamily::Normal)
.with_split_direction(DistSplitDirection::Cyclic),
))
.multi_strategy(MultiStrategy::MultiOutputTree)
.build()
.unwrap(),
matrix(1),
|m| {
matches!(m.objective().built_in(), Some(Objective::Dist(d))
if d.family() == DistFamily::Normal
&& d.split_direction() == Some(DistSplitDirection::Cyclic))
&& m.has_vector_leaves()
},
),
(
"dist:negbinomial",
base()
.objective(Objective::Dist(Distributional::new(
DistFamily::NegativeBinomial,
)))
.build()
.unwrap(),
negbinomial,
|m| {
m.objective().built_in().and_then(Objective::dist_family)
== Some(DistFamily::NegativeBinomial)
},
),
];
let mut models: Vec<(&str, BoostedModel, DMatrix, Has)> = cases
.into_iter()
.map(|(name, params, data, has)| (name, train(¶ms, &data, 4).unwrap(), data, has))
.collect();
let plain = train(&base().build().unwrap(), &matrix(1), 4).unwrap();
let mut doc: Value = serde_json::from_slice(&plain.encode(ModelFormat::Json).unwrap()).unwrap();
doc["best_iteration"] = 1.into();
let stopped = BoostedModel::decode(doc.to_string(), ModelFormat::Json).unwrap();
models.push(("early stopping", stopped, matrix(1), |m| {
m.best_iteration() == Some(1)
}));
models
}
#[test]
fn native_formats_round_trip_every_model_feature() {
for (name, model, data, has_feature) in feature_models() {
assert!(has_feature(&model), "{name}: feature not exercised");
let bytes = model.encode(ModelFormat::Binary).unwrap();
let from_binary = BoostedModel::decode(&bytes, ModelFormat::Binary).unwrap();
assert_eq!(
from_binary.encode(ModelFormat::Binary).unwrap(),
bytes,
"{name}: binary"
);
let from_json =
BoostedModel::decode(model.encode(ModelFormat::Json).unwrap(), ModelFormat::Json)
.unwrap();
assert_eq!(
from_json.encode(ModelFormat::Binary).unwrap(),
bytes,
"{name}: JSON"
);
let expected = bits(model.predict(&data, Iterations::Best).unwrap().as_slice());
for restored in [&from_binary, &from_json] {
assert_eq!(
bits(
restored
.predict(&data, Iterations::Best)
.unwrap()
.as_slice()
),
expected,
"{name}"
);
assert_eq!(restored.n_outputs(), model.n_outputs(), "{name}");
assert_eq!(restored.best_iteration(), model.best_iteration(), "{name}");
assert_eq!(restored.n_targets(), model.n_targets(), "{name}");
assert_eq!(
restored.num_parallel_tree(),
model.num_parallel_tree(),
"{name}"
);
assert_eq!(
restored.has_vector_leaves(),
model.has_vector_leaves(),
"{name}"
);
assert_eq!(restored.objective(), model.objective(), "{name}");
}
}
}
fn saved_dir(version: &str) -> PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests/data/saved")
.join(version)
}
fn slug(name: &str) -> String {
name.chars()
.map(|c| if c.is_ascii_alphanumeric() { c } else { '_' })
.collect()
}
fn margin_bytes(margins: &[f32]) -> Vec<u8> {
margins.iter().flat_map(|m| m.to_le_bytes()).collect()
}
#[test]
fn saved_models_keep_loading_with_their_margins() {
let root = saved_dir("");
let mut versions: Vec<PathBuf> = std::fs::read_dir(&root)
.unwrap()
.map(|entry| entry.unwrap().path())
.collect();
versions.sort();
assert!(
!versions.is_empty(),
"no saved models under {}",
root.display()
);
let cases = feature_models();
for dir in versions {
for (name, _, data, _) in &cases {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
let expected = std::fs::read(file("margins")).unwrap();
let binary = BoostedModel::load(file("bin"), ModelFormat::Binary).unwrap();
let json = BoostedModel::load(file("json"), ModelFormat::Json).unwrap();
for (format, model) in [("bin", binary), ("json", json)] {
let margins = margin_bytes(
model
.predict_margin(data, Iterations::Best)
.unwrap()
.as_slice(),
);
assert!(margins == expected, "{}: {name} ({format})", dir.display());
}
if file("hbtd").exists() {
let compact = CompactModel::load(file("hbtd")).unwrap();
let margins = margin_bytes(compact.predict_margin(data).unwrap().as_slice());
assert!(margins == expected, "{}: {name} (compact)", dir.display());
}
}
for (name, case) in diffusion_models() {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
if !file("hbdm").exists() {
continue;
}
let expected = std::fs::read(file("hbdm.probe")).unwrap();
let binary = DiffusionModel::load(file("hbdm"), DiffusionFormat::Binary).unwrap();
let json = DiffusionModel::load(file("hbdm.json"), DiffusionFormat::Json).unwrap();
for (format, model) in [("hbdm", binary), ("hbdm.json", json)] {
assert_eq!(model.method(), case.method(), "{name} ({format})");
let margins = regressor_margins(&model);
assert!(margins == expected, "{}: {name} ({format})", dir.display());
assert!(
model
.sample(&matrix(1), 2, &SampleOptions::seeded(0))
.is_ok(),
"{name} ({format})"
);
}
}
for (name, case) in forest_models() {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
if !file("hbff").exists() {
continue;
}
let expected = std::fs::read(file("hbff.probe")).unwrap();
let binary = ForestModel::load(file("hbff"), DiffusionFormat::Binary).unwrap();
let json = ForestModel::load(file("hbff.json"), DiffusionFormat::Json).unwrap();
for (format, model) in [("hbff", binary), ("hbff.json", json)] {
assert_eq!(model.method(), case.method(), "{name} ({format})");
assert_eq!(model.classes(), case.classes(), "{name} ({format})");
let margins = forest_margins(&model);
assert!(margins == expected, "{}: {name} ({format})", dir.display());
assert!(model.sample(2, 0).is_ok(), "{name} ({format})");
}
}
}
}
#[test]
fn saved_models_re_save_byte_identically() {
let mut versions = 0;
for dir in std::fs::read_dir(saved_dir("")).unwrap() {
let dir = dir.unwrap().path();
let current = dir.file_name().unwrap() == env!("CARGO_PKG_VERSION");
let mut checked = 0;
for entry in std::fs::read_dir(&dir).unwrap() {
let path = entry.unwrap().path();
let what = path.display();
let stored = std::fs::read(&path).unwrap();
match path.extension().and_then(|e| e.to_str()) {
Some("bin") => {
let json_path = path.with_extension("json");
let json = std::fs::read_to_string(&json_path).unwrap();
let saved_json: Value = serde_json::from_str(&json).unwrap();
let saved_json = saved_json.as_object().unwrap();
let hbtd = std::fs::read(path.with_extension("hbtd")).ok();
let from_bin = BoostedModel::decode(&stored, ModelFormat::Binary).unwrap();
let from_json = BoostedModel::decode(&json, ModelFormat::Json).unwrap();
for (source, model) in [("bin", from_bin), ("json", from_json)] {
if current {
let bytes = model.encode(ModelFormat::Binary).unwrap();
assert!(bytes == stored, "{what}: {source} re-saved as bin");
}
let resaved: Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap())
.unwrap();
let resaved = resaved.as_object().unwrap();
for key in saved_json.keys() {
assert!(resaved.contains_key(key), "{what}: {source} drops `{key}`");
}
for (key, value) in resaved {
let saved = saved_json.get(key).unwrap_or(&Value::Null);
assert!(value == saved, "{what}: {source} re-saves `{key}`");
}
if let Some(hbtd) = &hbtd {
let compact = model.to_compact_bytes().unwrap();
assert!(compact == *hbtd, "{what}: {source} re-saved as compact");
}
}
}
Some("hbdm") if current => {
let resaved = DiffusionModel::decode(&stored, DiffusionFormat::Binary)
.unwrap()
.encode(DiffusionFormat::Binary)
.unwrap();
assert!(resaved == stored, "{what} re-saves differently");
}
Some("hbff") if current => {
let resaved = ForestModel::decode(&stored, DiffusionFormat::Binary)
.unwrap()
.encode(DiffusionFormat::Binary)
.unwrap();
assert!(resaved == stored, "{what} re-saves differently");
}
_ => continue,
}
checked += 1;
}
assert!(checked > 0, "no saved models under {}", dir.display());
versions += 1;
}
assert!(versions > 0, "no saved model versions");
}
#[test]
#[ignore = "run once per release, then commit tests/data/saved/<version>"]
fn save_models_of_this_version() {
let dir = saved_dir(env!("CARGO_PKG_VERSION"));
assert!(
!dir.exists(),
"{} exists; saved models of a version are never rewritten",
dir.display()
);
std::fs::create_dir_all(&dir).unwrap();
for (name, model, data, _) in feature_models() {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
model.save(file("bin"), ModelFormat::Binary).unwrap();
model.save(file("json"), ModelFormat::Json).unwrap();
if let Ok(compact) = model.to_compact() {
compact.save(file("hbtd")).unwrap();
}
let margins = model.predict_margin(&data, Iterations::Best).unwrap();
std::fs::write(file("margins"), margin_bytes(margins.as_slice())).unwrap();
}
for (name, model) in diffusion_models() {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
model.save(file("hbdm"), DiffusionFormat::Binary).unwrap();
model
.save(file("hbdm.json"), DiffusionFormat::Json)
.unwrap();
std::fs::write(file("hbdm.probe"), regressor_margins(&model)).unwrap();
}
for (name, model) in forest_models() {
let file = |ext: &str| dir.join(format!("{}.{ext}", slug(name)));
model.save(file("hbff"), DiffusionFormat::Binary).unwrap();
model
.save(file("hbff.json"), DiffusionFormat::Json)
.unwrap();
std::fs::write(file("hbff.probe"), forest_margins(&model)).unwrap();
}
}
fn diffusion_models() -> Vec<(&'static str, DiffusionModel)> {
let tiny = |mut params: DiffusionParams| {
params.n_repeats = std::num::NonZeroUsize::new(2).unwrap();
params.num_boost_round = std::num::NonZeroUsize::new(4).unwrap();
params.early_stopping = None;
params.training.nthread = std::num::NonZeroUsize::new(1);
if let Some(r) = &mut params.residualizer {
r.num_boost_round = std::num::NonZeroUsize::new(3).unwrap();
}
params
};
let mut vp = ScoreConfig::treeffuser();
vp.sde = Sde::VariancePreserving {
beta_min: 0.1,
beta_max: 20.0,
};
let mut treeffuser_vp = tiny(DiffusionParams::treeffuser());
treeffuser_vp.method = Method::Score(vp);
[
("diffusion score", tiny(DiffusionParams::default()), 1),
("diffusion treeffuser vp", treeffuser_vp, 2),
(
"diffusion flow matching",
tiny(DiffusionParams::flow_matching()),
1,
),
]
.into_iter()
.map(|(name, params, k)| (name, DiffusionModel::fit(¶ms, &matrix(k)).unwrap()))
.collect()
}
fn regressor_margins(model: &DiffusionModel) -> Vec<u8> {
let regressor = model.regressor();
let cols = regressor.n_features();
let x: Vec<f32> = (0..16 * cols)
.map(|i| ((i * 37) % 101) as f32 / 25.0 - 2.0)
.collect();
let probe = DMatrix::from_dense(&x, 16, cols).unwrap();
margin_bytes(
regressor
.predict_margin(&probe, Iterations::Best)
.unwrap()
.as_slice(),
)
}
fn forest_models() -> Vec<(&'static str, ForestModel)> {
let tiny = |mut params: ForestParams, kinds: Option<Vec<ColumnKind>>| {
params.n_t = NoiseLevels::new(3).unwrap();
params.duplicate_k = std::num::NonZeroUsize::new(2).unwrap();
params.num_boost_round = std::num::NonZeroUsize::new(3).unwrap();
params.training.nthread = std::num::NonZeroUsize::new(1);
params.column_kinds = kinds;
params
};
let n = 160;
let complete: Vec<f32> = (0..n)
.flat_map(|i| {
let [a, b, c, _] = four_features(i);
[a, b, c, (i % 3) as f32]
})
.collect();
let with_missing: Vec<f32> = (0..n).flat_map(four_features).collect();
let classes: Vec<f32> = (0..n).map(|i| (i % 2) as f32).collect();
let mut kinds = vec![ColumnKind::Continuous; COLS];
kinds[3] = ColumnKind::Categorical;
let flow = ForestModel::fit(
&tiny(ForestParams::default(), Some(kinds)),
&DMatrix::from_dense(&complete, n, COLS).unwrap(),
)
.unwrap();
let diffusion = ForestModel::fit(
&tiny(ForestParams::forest_diffusion(), None),
&DMatrix::from_dense(&with_missing, n, COLS)
.unwrap()
.with_labels(&classes)
.unwrap(),
)
.unwrap();
vec![("forest flow", flow), ("forest diffusion", diffusion)]
}
fn forest_margins(model: &ForestModel) -> Vec<u8> {
let mut out = Vec::new();
for gbdt in model.gbdts() {
let cols = gbdt.n_features();
let x: Vec<f32> = (0..8 * cols)
.map(|i| ((i * 37) % 101) as f32 / 50.0 - 1.0)
.collect();
let probe = DMatrix::from_dense(&x, 8, cols).unwrap();
out.extend(margin_bytes(
gbdt.predict_margin(&probe, Iterations::Best)
.unwrap()
.as_slice(),
));
}
out
}