use hessboost::config::{ProcessType, Refresh};
use hessboost::data::FeatureType;
use hessboost::internals::HistCuts;
use hessboost::model::Iterations;
use hessboost::model::ModelFormat;
use hessboost::prelude::{BoostedModel, DMatrix, HessboostError, Trainer, TrainingParams, train};
use hessboost::training::EvalHistory;
use serde::{Deserialize, Deserializer};
use serde_json::{Map, Value};
use std::cmp::Ordering;
use std::collections::BTreeMap;
use std::path::Path;
mod common;
use common::fixtures::{fixtures_dir, load_all, nan_for_null};
const CONTRIB_ROWS: usize = 50;
const INTERACTION_ROWS: usize = 5;
#[derive(Deserialize, Clone, Copy, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
enum Tier {
Exact,
Quality,
Trainonly,
}
#[derive(Deserialize)]
struct Tol {
train: f64,
import: f64,
contribs: f64,
evals: f64,
interactions: f64,
}
#[derive(Deserialize)]
struct Fixture {
name: String,
tier: Tier,
params: Map<String, Value>,
num_class: usize,
num_round: usize,
n_train: usize,
n_test: usize,
n_cols: usize,
n_targets: usize,
#[serde(deserialize_with = "nan_for_null")]
x_train: Vec<f32>,
y_train: Vec<f32>,
#[serde(deserialize_with = "nan_for_null")]
x_test: Vec<f32>,
y_test: Vec<f32>,
#[serde(default, deserialize_with = "bounds")]
label_lower_bound: Option<Vec<f32>>,
#[serde(default, deserialize_with = "bounds")]
label_upper_bound: Option<Vec<f32>>,
#[serde(default, deserialize_with = "bounds")]
test_label_lower_bound: Option<Vec<f32>>,
#[serde(default, deserialize_with = "bounds")]
test_label_upper_bound: Option<Vec<f32>>,
weights: Option<Vec<f32>>,
#[serde(default)]
feature_weights: Option<Vec<f32>>,
#[serde(default)]
test_weights: Option<Vec<f32>>,
group_sizes: Option<Vec<usize>>,
test_group_sizes: Option<Vec<usize>>,
feature_types: Option<Vec<String>>,
xgb_pred: Vec<f32>,
xgb_margin: Vec<f32>,
xgb_contribs: Vec<f32>,
#[serde(default)]
xgb_evals: Option<BTreeMap<String, Vec<f64>>>,
xgb_interactions: Option<Vec<f32>>,
xgb_model: Value,
xgb_model_ubj: String,
tol: Tol,
continuation: Option<Continuation>,
refresh: Option<RefreshCase>,
#[serde(default)]
ranges: Vec<RangeCase>,
#[serde(default)]
range_contribs: Vec<RangeContribs>,
#[serde(default)]
slices: Vec<SliceCase>,
}
#[derive(Deserialize)]
struct Continuation {
first_rounds: usize,
xgb_model_initial: Value,
}
#[derive(Deserialize)]
struct RefreshCase {
n_rows: usize,
y: Vec<f32>,
rounds: usize,
refresh_leaf: bool,
xgb_pred: Vec<f32>,
}
#[derive(Deserialize)]
struct RangeCase {
begin: usize,
end: usize,
margin: Vec<f32>,
}
#[derive(Deserialize)]
struct RangeContribs {
end: usize,
contribs: Vec<f32>,
leaf: Vec<f32>,
}
#[derive(Deserialize)]
struct SliceCase {
begin: usize,
end: usize,
step: usize,
margin: Vec<f32>,
}
#[derive(Deserialize)]
struct CutFixture {
name: String,
tree_method: String,
objective: String,
max_bin: usize,
n_rows: usize,
n_cols: usize,
#[serde(deserialize_with = "nan_for_null")]
x: Vec<f32>,
w: Option<Vec<f32>>,
indptr: Vec<usize>,
#[serde(deserialize_with = "nan_for_null")]
cuts: Vec<f32>,
}
fn bounds<'de, D: Deserializer<'de>>(d: D) -> Result<Option<Vec<f32>>, D::Error> {
#[derive(Deserialize)]
#[serde(untagged)]
enum Bound {
Number(f32),
Text(String),
}
let Some(v) = Option::<Vec<Bound>>::deserialize(d)? else {
return Ok(None);
};
v.into_iter()
.map(|b| match b {
Bound::Number(x) => Ok(x),
Bound::Text(s) if s == "inf" => Ok(f32::INFINITY),
Bound::Text(s) if s == "-inf" => Ok(f32::NEG_INFINITY),
Bound::Text(s) => Err(serde::de::Error::custom(format!("invalid bound `{s}`"))),
})
.collect::<Result<Vec<_>, _>>()
.map(Some)
}
fn build_params(fx: &Fixture) -> Result<TrainingParams, String> {
if let Some(v) = fx.params.get("lambdarank_num_pair_per_sample") {
let n = v.as_u64().ok_or_else(|| {
format!("`lambdarank_num_pair_per_sample` must be an integer, got {v}")
})?;
let groups = fx
.group_sizes
.as_deref()
.ok_or("lambdarank_num_pair_per_sample without group_sizes")?;
if groups.iter().any(|&g| (g as u64) < n) {
return Err(format!(
"`lambdarank_num_pair_per_sample`={n} exceeds a train group size"
));
}
}
TrainingParams::from_xgboost(fx.params.clone()).map_err(|e| format!("invalid params: {e}"))
}
fn xgb_range(model: &BoostedModel, begin: usize, end: usize) -> std::ops::Range<usize> {
begin..if end == 0 {
model.num_boost_rounds()
} else {
end
}
}
fn max_abs_diff(what: &str, a: &[f32], b: &[f32]) -> Result<f64, String> {
if a.len() != b.len() {
return Err(format!(
"{what}: length mismatch (hessboost {}, xgboost {})",
a.len(),
b.len()
));
}
Ok(a.iter().zip(b).fold(0.0f64, |m, (x, y)| {
let d = (f64::from(*x) - f64::from(*y)).abs();
m.max(if d.is_nan() { f64::INFINITY } else { d })
}))
}
fn rmse(p: &[f32], y: &[f32]) -> f64 {
let s: f64 = p
.iter()
.zip(y)
.map(|(a, b)| (f64::from(*a) - f64::from(*b)).powi(2))
.sum();
(s / y.len() as f64).sqrt()
}
fn accuracy(objective: &str, p: &[f32], y: &[f32], k: usize) -> f64 {
let hits = y
.iter()
.enumerate()
.filter(|&(i, &yi)| {
let class = if objective == "multi:softprob" {
let row = &p[i * k..(i + 1) * k];
row.iter()
.enumerate()
.fold(0usize, |best, (c, v)| if *v > row[best] { c } else { best })
} else if objective.starts_with("binary:") {
usize::from(p[i] > 0.5)
} else {
p[i].round() as usize
};
class == yi as usize
})
.count();
hits as f64 / y.len() as f64
}
fn dcg(labels_in_rank_order: impl Iterator<Item = f32>) -> f64 {
labels_in_rank_order
.enumerate()
.map(|(i, l)| (2f64.powf(f64::from(l)) - 1.0) / (i as f64 + 2.0).log2())
.sum()
}
fn mean_ndcg(scores: &[f32], labels: &[f32], groups: &[usize]) -> f64 {
let mut start = 0;
let mut total = 0.0;
for &g in groups {
let s = &scores[start..start + g];
let l = &labels[start..start + g];
let mut order: Vec<usize> = (0..g).collect();
order.sort_by(|&a, &b| {
s[b].partial_cmp(&s[a])
.unwrap_or(Ordering::Equal)
.then(a.cmp(&b))
});
let mut ideal = l.to_vec();
ideal.sort_by(|a, b| b.partial_cmp(a).unwrap_or(Ordering::Equal));
let idcg = dcg(ideal.iter().copied());
if idcg > 0.0 {
total += dcg(order.iter().map(|&j| l[j])) / idcg;
}
start += g;
}
total / groups.len() as f64
}
fn fmt_delta(d: &Result<f64, String>) -> String {
match d {
Ok(v) => format!("{v:.2e}"),
Err(_) => "ERR".to_string(),
}
}
struct Row {
name: String,
tier: &'static str,
train: String,
import: String,
margin: String,
contribs: String,
ubj: String,
interactions: String,
export: String,
evals: String,
extra: String,
base_score: String,
}
struct Case<'a> {
fx: &'a Fixture,
failures: &'a mut Vec<String>,
}
impl Case<'_> {
fn fail(&mut self, msg: impl std::fmt::Display) {
self.failures.push(format!("{}: {msg}", self.fx.name));
}
fn check(&mut self, what: &str, delta: &Result<f64, String>, tol: f64) -> String {
match delta {
Ok(d) if *d <= tol => {}
Ok(d) => self.fail(format!("{what}: max|Δ|={d:.3e} > tol {tol:.0e}")),
Err(e) => self.fail(format!("{what}: {e}")),
}
fmt_delta(delta)
}
fn dmatrix(&self, x: &[f32], n_rows: usize) -> Result<DMatrix, String> {
let mut d =
DMatrix::from_dense(x, n_rows, self.fx.n_cols).map_err(|e| format!("DMatrix: {e}"))?;
if let Some(types) = &self.fx.feature_types {
let types: Vec<FeatureType> = types
.iter()
.map(|t| match t.as_str() {
"c" => Ok(FeatureType::Categorical),
"q" => Ok(FeatureType::Numerical),
other => Err(format!("unsupported feature type `{other}`")),
})
.collect::<Result<_, _>>()?;
d = d
.with_feature_types(&types)
.map_err(|e| format!("feature types: {e}"))?;
}
Ok(d)
}
fn with_meta(
&self,
d: DMatrix,
labels: &[f32],
bounds: (Option<&Vec<f32>>, Option<&Vec<f32>>),
weights: Option<&Vec<f32>>,
groups: Option<&Vec<usize>>,
feature_weights: Option<&Vec<f32>>,
) -> Result<DMatrix, String> {
let mut d = d;
if !labels.is_empty() {
d = d
.with_label_matrix(labels, self.fx.n_targets)
.map_err(|e| format!("labels: {e}"))?;
}
match bounds {
(Some(lo), Some(hi)) => {
d = d
.with_label_bounds(lo, hi)
.map_err(|e| format!("label bounds: {e}"))?;
}
(None, None) => {}
_ => return Err("only one of the label bounds is present".to_string()),
}
if let Some(w) = weights {
d = d.with_weights(w).map_err(|e| format!("weights: {e}"))?;
}
if let Some(w) = feature_weights {
d = d
.with_feature_weights(w)
.map_err(|e| format!("feature_weights: {e}"))?;
}
if let Some(g) = groups {
d = d
.with_group_sizes(g)
.map_err(|e| format!("group_sizes: {e}"))?;
}
Ok(d)
}
fn train_matrix(&self) -> Result<DMatrix, String> {
let fx = self.fx;
self.with_meta(
self.dmatrix(&fx.x_train, fx.n_train)?,
&fx.y_train,
(fx.label_lower_bound.as_ref(), fx.label_upper_bound.as_ref()),
fx.weights.as_ref(),
fx.group_sizes.as_ref(),
fx.feature_weights.as_ref(),
)
}
fn eval_matrix(&self) -> Result<DMatrix, String> {
let fx = self.fx;
self.with_meta(
self.dmatrix(&fx.x_test, fx.n_test)?,
&fx.y_test,
(
fx.test_label_lower_bound.as_ref(),
fx.test_label_upper_bound.as_ref(),
),
fx.test_weights.as_ref(),
fx.test_group_sizes.as_ref(),
None,
)
}
fn train_and_compare(
&mut self,
dtest: &DMatrix,
) -> Result<(BoostedModel, Vec<f32>, EvalHistory), String> {
let fx = self.fx;
let params = build_params(fx)?;
let dtrain = self.train_matrix()?;
let (model, history) = match &fx.continuation {
Some(c) => {
let first =
train(¶ms, &dtrain, c.first_rounds).map_err(|e| format!("train: {e}"))?;
let model = Trainer::new(¶ms, &dtrain, fx.num_round - c.first_rounds)
.init_model(&first)
.train()
.map(|r| r.model)
.map_err(|e| format!("continue training: {e}"))?;
(model, EvalHistory::default())
}
None if fx.xgb_evals.is_some() => {
let deval = self.eval_matrix()?;
let result = Trainer::new(¶ms, &dtrain, fx.num_round)
.eval(&deval, "test")
.train()
.map_err(|e| format!("train: {e}"))?;
(result.model, result.history)
}
None => {
let model =
train(¶ms, &dtrain, fx.num_round).map_err(|e| format!("train: {e}"))?;
(model, EvalHistory::default())
}
};
let preds = model
.predict(dtest, Iterations::Best)
.map_err(|e| format!("predict: {e}"))?;
if preds.as_slice().len() != fx.xgb_pred.len() {
return Err(format!(
"train predict: length mismatch (hessboost {}, xgboost {})",
preds.as_slice().len(),
fx.xgb_pred.len()
));
}
Ok((model, preds.into_vec(), history))
}
fn compare_evals(&mut self, history: &EvalHistory) -> String {
let fx = self.fx;
let Some(oracle) = &fx.xgb_evals else {
return "-".to_string();
};
let mut ours: BTreeMap<&str, Vec<f64>> = BTreeMap::new();
for round in history.rounds() {
for (_, metric, value) in round.scores() {
ours.entry(metric).or_default().push(value);
}
}
let our_names: Vec<&str> = ours.keys().copied().collect();
let xgb_names: Vec<&str> = oracle.keys().map(String::as_str).collect();
if our_names != xgb_names {
self.fail(format!(
"evals: metric names {our_names:?} != xgboost {xgb_names:?}"
));
return "ERR".to_string();
}
let mut worst = 0.0f64;
for (name, want) in oracle {
let got = &ours[name.as_str()];
if got.len() != want.len() {
self.fail(format!(
"evals {name}: {} rounds, xgboost {}",
got.len(),
want.len()
));
return "ERR".to_string();
}
for (round, (a, b)) in got.iter().zip(want).enumerate() {
let rel = (a - b).abs() / b.abs().max(1.0);
if rel.is_nan() || rel > fx.tol.evals {
self.fail(format!(
"evals {name} round {round}: hessboost {a} xgboost {b} (rel {rel:.3e} > tol {:.0e})",
fx.tol.evals
));
return "ERR".to_string();
}
worst = worst.max(rel);
}
}
format!("{worst:.2e}")
}
fn continue_imported(&self, dtest: &DMatrix, c: &Continuation) -> Result<f64, String> {
let fx = self.fx;
let params = build_params(fx)?;
let initial =
BoostedModel::decode(c.xgb_model_initial.to_string(), ModelFormat::XgboostJson)
.map_err(|e| format!("import initial model: {e}"))?;
let model = Trainer::new(
¶ms,
&self.train_matrix()?,
fx.num_round - c.first_rounds,
)
.init_model(&initial)
.train()
.map(|r| r.model)
.map_err(|e| format!("continue training: {e}"))?;
let preds = model
.predict(dtest, Iterations::Best)
.map_err(|e| e.to_string())?;
max_abs_diff("continue imported", preds.as_slice(), &fx.xgb_pred)
}
fn refresh(
&self,
base: &BoostedModel,
dtest: &DMatrix,
r: &RefreshCase,
) -> Result<f64, String> {
let fx = self.fx;
let mut params = build_params(fx)?;
params.process_type = ProcessType::Update(if r.refresh_leaf {
Refresh::default()
} else {
Refresh::stats_only()
});
let data = self
.dmatrix(&fx.x_train[..r.n_rows * fx.n_cols], r.n_rows)?
.with_labels(&r.y)
.map_err(|e| format!("refresh labels: {e}"))?;
let model = Trainer::new(¶ms, &data, r.rounds)
.init_model(base)
.train()
.map(|r| r.model)
.map_err(|e| format!("refresh: {e}"))?;
let preds = model
.predict(dtest, Iterations::Best)
.map_err(|e| e.to_string())?;
max_abs_diff("refresh", preds.as_slice(), &r.xgb_pred)
}
fn range_checks(
&self,
model: &BoostedModel,
dtest: &DMatrix,
dcontrib: &DMatrix,
tol: f64,
with_leaves: bool,
) -> Vec<(String, Result<f64, String>, f64)> {
let fx = self.fx;
let mut out = Vec::new();
for r in &fx.ranges {
let what = format!("margin range [{}, {})", r.begin, r.end);
let d = model
.predict_margin(dtest, xgb_range(model, r.begin, r.end))
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff(&what, p.as_slice(), &r.margin));
out.push((what, d, tol));
}
for r in &fx.range_contribs {
let what = format!("contribs range [0, {})", r.end);
let d = model
.predict_contribs(dcontrib, xgb_range(model, 0, r.end))
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff(&what, p.as_slice(), &r.contribs));
out.push((what, d, fx.tol.contribs));
if with_leaves {
let what = format!("leaf range [0, {})", r.end);
let d = model
.predict_leaf(dcontrib, xgb_range(model, 0, r.end))
.map_err(|e| e.to_string())
.and_then(|p| {
let p: Vec<f32> = p.as_slice().iter().map(|&l| l as f32).collect();
max_abs_diff(&what, &p, &r.leaf)
});
out.push((what, d, 0.0));
}
}
for s in &fx.slices {
let what = format!("slice [{}:{}:{}]", s.begin, s.end, s.step);
let d = model
.slice(xgb_range(model, s.begin, s.end), s.step)
.and_then(|m| m.predict_margin(dtest, Iterations::Best))
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff(&what, p.as_slice(), &s.margin));
out.push((what, d, tol));
}
out
}
fn extras(
&mut self,
imported: &hessboost::error::Result<BoostedModel>,
trained: Option<&BoostedModel>,
dtest: &DMatrix,
dcontrib: &DMatrix,
) -> String {
let fx = self.fx;
if fx.continuation.is_none()
&& fx.refresh.is_none()
&& fx.ranges.is_empty()
&& fx.range_contribs.is_empty()
&& fx.slices.is_empty()
{
return "-".to_string();
}
let imported = match imported {
Ok(m) => m,
Err(e) => {
self.fail(format!("extras import: {e}"));
return "ERR".to_string();
}
};
let trained = trained.filter(|_| fx.tier == Tier::Exact);
let mut checks = Vec::new();
if let Some(c) = &fx.continuation {
checks.push((
"continue imported".to_string(),
self.continue_imported(dtest, c),
fx.tol.train,
));
}
if let Some(r) = &fx.refresh {
checks.push((
"refresh imported".to_string(),
self.refresh(imported, dtest, r),
fx.tol.train,
));
if let Some(t) = trained {
checks.push((
"refresh trained".to_string(),
self.refresh(t, dtest, r),
fx.tol.train,
));
}
}
checks.extend(self.range_checks(imported, dtest, dcontrib, fx.tol.import, true));
if let Some(t) = trained {
checks.extend(self.range_checks(t, dtest, dcontrib, fx.tol.train, false));
}
let mut worst = 0.0f64;
for (what, delta, tol) in &checks {
self.check(what, delta, *tol);
worst = worst.max(*delta.as_ref().unwrap_or(&f64::INFINITY));
}
format!("{worst:.1e}/{}", checks.len())
}
fn quality_band(&mut self, preds: &[f32]) -> String {
let fx = self.fx;
let objective = fx.params["objective"].as_str().unwrap_or("");
let band = fx.tol.train;
let (metric, seq, xgb, ok) = if objective.starts_with("rank:") {
let Some(groups) = fx.test_group_sizes.as_deref() else {
self.fail("quality band for rank:* needs test_group_sizes");
return "ERR".to_string();
};
let s = mean_ndcg(preds, &fx.y_test, groups);
let x = mean_ndcg(&fx.xgb_pred, &fx.y_test, groups);
("ndcg", s, x, s >= x - band)
} else if objective.starts_with("binary:") || objective.starts_with("multi:") {
let s = accuracy(objective, preds, &fx.y_test, fx.num_class);
let x = accuracy(objective, &fx.xgb_pred, &fx.y_test, fx.num_class);
("acc", s, x, s >= x - band)
} else {
let s = rmse(preds, &fx.y_test);
let x = rmse(&fx.xgb_pred, &fx.y_test);
("rmse", s, x, s <= x * band + 1e-6)
};
if !ok {
self.fail(format!(
"quality band: {metric} hessboost={seq:.5} xgboost={xgb:.5} (band {band})"
));
}
format!("{metric} {seq:.4}/{xgb:.4}")
}
fn import_and_compare(
&mut self,
imported: &hessboost::error::Result<BoostedModel>,
dtest: &DMatrix,
dcontrib: &DMatrix,
) -> [String; 4] {
let fx = self.fx;
if booster_of(fx) == "gblinear" {
return match imported {
Err(HessboostError::ModelFormat(_)) => std::array::from_fn(|_| "n/a".to_string()),
Err(e) => {
self.fail(format!("expected ModelFormat import error, got {e}"));
std::array::from_fn(|_| "ERR".to_string())
}
Ok(_) => {
self.fail("unsupported gblinear import unexpectedly succeeded");
std::array::from_fn(|_| "ERR".to_string())
}
};
}
let model = match imported {
Ok(m) => m,
Err(e) => {
self.fail(format!("import: {e}"));
return std::array::from_fn(|_| "ERR".to_string());
}
};
let pred = model
.predict(dtest, Iterations::Best)
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff("import predict", p.as_slice(), &fx.xgb_pred));
let margin = model
.predict_margin(dtest, Iterations::Best)
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff("import margin", p.as_slice(), &fx.xgb_margin));
let contribs = model
.predict_contribs(dcontrib, Iterations::Best)
.map_err(|e| e.to_string())
.and_then(|p| max_abs_diff("import contribs", p.as_slice(), &fx.xgb_contribs));
let interactions = match &fx.xgb_interactions {
None => "-".to_string(),
Some(want) => {
let delta = self
.dmatrix(&fx.x_test[..INTERACTION_ROWS * fx.n_cols], INTERACTION_ROWS)
.and_then(|d| {
model
.predict_interactions(&d, Iterations::Best)
.map_err(|e| e.to_string())
})
.and_then(|p| max_abs_diff("import interactions", p.as_slice(), want));
self.check("import interactions", &delta, fx.tol.interactions)
}
};
[
self.check("import predict", &pred, fx.tol.import),
self.check("import margin", &margin, fx.tol.import),
self.check("import contribs", &contribs, fx.tol.contribs),
interactions,
]
}
fn import_ubjson(
&mut self,
from_json: &hessboost::error::Result<BoostedModel>,
dir: &Path,
) -> String {
let fx = self.fx;
let from_ubj = std::fs::read(dir.join(&fx.xgb_model_ubj))
.map_err(HessboostError::from)
.and_then(|bytes| BoostedModel::decode(&bytes, ModelFormat::XgboostUbjson));
let verdict = match (&from_ubj, from_json) {
(Ok(u), Ok(j)) => {
match (u.encode(ModelFormat::Binary), j.encode(ModelFormat::Binary)) {
(Ok(u), Ok(j)) if u == j => Ok("same"),
(Ok(_), Ok(_)) => Err("UBJSON import differs from the JSON import".to_string()),
(Err(e), _) | (_, Err(e)) => Err(format!("encode imported model: {e}")),
}
}
(Err(HessboostError::ModelFormat(_)), Err(HessboostError::ModelFormat(_))) => Ok("n/a"),
(Err(e), _) => Err(format!("UBJSON import: {e}")),
(Ok(_), Err(e)) => Err(format!("UBJSON import succeeded, JSON import failed: {e}")),
};
match verdict {
Ok(cell) => cell.to_string(),
Err(e) => {
self.fail(e);
"ERR".to_string()
}
}
}
fn export(&mut self, model: &BoostedModel, preds: &[f32], dir: &Path) -> String {
let fx = self.fx;
if booster_of(fx) == "gblinear" {
return "skipped".to_string();
}
let written = model
.encode(ModelFormat::XgboostJson)
.map_err(|e| format!("export: {e}"))
.and_then(|json| {
std::fs::write(dir.join(format!("{}.model.json", fx.name)), json)
.map_err(|e| format!("write model: {e}"))
})
.and_then(|()| {
model
.encode(ModelFormat::XgboostUbjson)
.map_err(|e| format!("export UBJSON: {e}"))
})
.and_then(|ubj| {
std::fs::write(dir.join(format!("{}.model.ubj", fx.name)), ubj)
.map_err(|e| format!("write UBJSON model: {e}"))
})
.and_then(|()| {
serde_json::to_string(preds)
.map_err(|e| format!("encode preds: {e}"))
.and_then(|p| {
std::fs::write(dir.join(format!("{}.pred.json", fx.name)), p)
.map_err(|e| format!("write preds: {e}"))
})
});
match written {
Ok(()) => "written".to_string(),
Err(e) => {
self.fail(e);
"ERR".to_string()
}
}
}
fn run(mut self, dir: &Path, exports: &Path) -> Row {
let fx = self.fx;
let tier = match fx.tier {
Tier::Exact => "exact",
Tier::Quality => "quality",
Tier::Trainonly => "train-only",
};
let mut row = Row {
name: fx.name.clone(),
tier,
train: "ERR".to_string(),
import: "ERR".to_string(),
margin: "ERR".to_string(),
contribs: "ERR".to_string(),
ubj: "ERR".to_string(),
interactions: "ERR".to_string(),
export: "n/a".to_string(),
evals: "-".to_string(),
extra: "-".to_string(),
base_score: "-".to_string(),
};
let dtest = match self.dmatrix(&fx.x_test, fx.n_test) {
Ok(d) => d,
Err(e) => {
self.fail(e);
return row;
}
};
let dcontrib = match self.dmatrix(&fx.x_test[..CONTRIB_ROWS * fx.n_cols], CONTRIB_ROWS) {
Ok(d) => d,
Err(e) => {
self.fail(e);
return row;
}
};
let trained = match self.train_and_compare(&dtest) {
Ok((model, preds, history)) => {
row.train = match fx.tier {
Tier::Exact | Tier::Trainonly => {
let d = max_abs_diff("train predict", &preds, &fx.xgb_pred);
self.check("train predict", &d, fx.tol.train)
}
Tier::Quality => self.quality_band(&preds),
};
row.evals = self.compare_evals(&history);
row.base_score = format!(
"{:?}",
model
.base_scores()
.iter()
.map(|v| (f64::from(*v) * 1e4).round() / 1e4)
.collect::<Vec<_>>()
);
row.export = self.export(&model, &preds, exports);
Some(model)
}
Err(e) => {
self.fail(e);
None
}
};
let imported = BoostedModel::decode(fx.xgb_model.to_string(), ModelFormat::XgboostJson);
[row.import, row.margin, row.contribs, row.interactions] =
self.import_and_compare(&imported, &dtest, &dcontrib);
row.ubj = self.import_ubjson(&imported, dir);
row.extra = self.extras(&imported, trained.as_ref(), &dtest, &dcontrib);
row
}
}
fn booster_of(fx: &Fixture) -> &str {
fx.params
.get("booster")
.and_then(Value::as_str)
.unwrap_or("gbtree")
}
fn xgb_base_score(fx: &Fixture) -> String {
fx.xgb_model
.pointer("/learner/learner_model_param/base_score")
.and_then(Value::as_str)
.unwrap_or("?")
.to_string()
}
#[test]
#[ignore = "requires fixtures from scripts/gen_fixtures.py"]
fn xgboost_parity() {
let dir = fixtures_dir();
let exports = dir.join("exports");
std::fs::create_dir_all(&exports).expect("create fixtures/exports");
let fixtures: Vec<Fixture> = load_all(&dir, "parity", "scripts/gen_fixtures.py");
let mut failures = Vec::new();
println!(
"{:<30} {:<7} {:<22} {:<9} {:<9} {:<9} {:<9} {:<5} {:<8} {:<9} {:<11} base_score hessboost | xgboost",
"case",
"tier",
"train",
"import",
"margin",
"contribs",
"inter",
"ubj",
"export",
"evals",
"extra"
);
for fx in &fixtures {
let row = Case {
fx,
failures: &mut failures,
}
.run(&dir, &exports);
println!(
"{:<30} {:<7} {:<22} {:<9} {:<9} {:<9} {:<9} {:<5} {:<8} {:<9} {:<11} {} | {}",
row.name,
row.tier,
row.train,
row.import,
row.margin,
row.contribs,
row.interactions,
row.ubj,
row.export,
row.evals,
row.extra,
row.base_score,
xgb_base_score(fx)
);
}
assert!(
failures.is_empty(),
"{} parity failure(s):\n {}",
failures.len(),
failures.join("\n ")
);
}
#[test]
#[ignore = "requires fixtures from scripts/gen_fixtures.py"]
fn quantile_cuts_match_xgboost() {
let fixtures: Vec<CutFixture> = load_all(
&fixtures_dir().join("cuts"),
"cut",
"scripts/gen_fixtures.py",
);
let mut failures: Vec<String> = Vec::new();
for fx in &fixtures {
let mut d = DMatrix::from_dense(&fx.x, fx.n_rows, fx.n_cols).expect("dense matrix");
if let Some(w) = &fx.w {
d = d.with_weights(w).expect("weights");
}
let cuts = match fx.tree_method.as_str() {
"hist" => HistCuts::from_dmatrix(&d, fx.max_bin),
"approx" => {
let unit = match fx.objective.as_str() {
"reg:squarederror" => 1.0,
"binary:logistic" => 0.25,
other => panic!("{}: unsupported cut oracle objective {other}", fx.name),
};
let hessians: Vec<f32> = (0..fx.n_rows)
.map(|r| unit * fx.w.as_ref().map_or(1.0, |w| w[r]))
.collect();
let const_hess = fx.objective == "reg:squarederror";
HistCuts::from_dmatrix_weighted(&d, fx.max_bin, &hessians, !const_hess)
}
other => panic!("{}: unsupported cut oracle tree_method {other}", fx.name),
};
let mut case_ok = true;
for f in 0..fx.n_cols {
let xgb = &fx.cuts[fx.indptr[f] + 1..fx.indptr[f + 1]];
let (start, end) = cuts.feature_bins(f);
let seq: Vec<f32> = (start..end).map(|i| cuts.cut_value(i)).collect();
let first_diff = seq
.iter()
.zip(xgb)
.position(|(a, b)| a.to_bits() != b.to_bits());
let problem = match first_diff {
Some(i) => Some(format!(
"feature {f} cut {i}: hessboost {:e} ({:#010x}) vs xgboost {:e} ({:#010x})",
seq[i],
seq[i].to_bits(),
xgb[i],
xgb[i].to_bits()
)),
None if seq.len() != xgb.len() => Some(format!(
"feature {f}: {} cuts vs xgboost {}",
seq.len(),
xgb.len()
)),
None => None,
};
if let Some(p) = problem {
case_ok = false;
failures.push(format!("{}: {p}", fx.name));
}
}
println!(
"{:<32} rows={:<7} cols={} max_bin={:<4} cuts={:<5} {}",
fx.name,
fx.n_rows,
fx.n_cols,
fx.max_bin,
cuts.total_bins(),
if case_ok { "OK" } else { "FAIL" }
);
}
assert!(
failures.is_empty(),
"{} cut mismatch(es):\n {}",
failures.len(),
failures.join("\n ")
);
}