use crate::config::ObjectiveParams;
use crate::error::{HessboostError, Result};
use crate::learner::BoostedModel;
use crate::learner::model::ModelSpec;
use crate::objective::{Objective, create_objective};
use crate::tree::{Node, RegTree};
use serde_json::{Map, Value, json};
const INVALID_NODE: i32 = i32::MAX;
pub fn export_xgboost_json(model: &BoostedModel) -> Result<String> {
if model.linear().is_some() {
return Err(HessboostError::model_format(
"XGBoost JSON export does not support gblinear models",
));
}
let num_feature = model.n_features();
let num_class = model.num_class();
let objective = model.objective().to_string();
let objective_impl = model.rebuild_objective().map_err(|_| {
HessboostError::model_format(format!(
"objective `{objective}` has no XGBoost equivalent; cannot export"
))
})?;
let n_outputs = model.n_outputs();
let n_trees = model.effective_ntrees();
let trees: Vec<Value> = model.trees()[..n_trees]
.iter()
.enumerate()
.map(|(id, t)| tree_to_json(id, t, num_feature))
.collect();
let tree_info: Vec<Value> = (0..n_trees)
.map(|t| json!((t % n_outputs) as i32))
.collect();
let mut booster_model = json!({
"gbtree_model_param": {
"num_parallel_tree": "1",
"num_trees": n_trees.to_string(),
},
"tree_info": tree_info,
"trees": trees,
});
if model.has_non_unit_tree_weights() {
let weight_drop: Vec<Value> = (0..n_trees).map(|t| json!(model.tree_weight(t))).collect();
booster_model["weight_drop"] = Value::Array(weight_drop);
}
let base_score = format_base_score(model.base_scores(), &*objective_impl);
let value = json!({
"version": [3, 4, 1],
"learner": {
"attributes": {},
"feature_names": [],
"feature_types": [],
"gradient_booster": {
"name": "gbtree",
"model": booster_model,
},
"learner_model_param": {
"base_score": base_score,
"boost_from_average": "0",
"num_class": num_class.to_string(),
"num_feature": num_feature.to_string(),
"num_target": "1",
},
"objective": objective_to_json(&objective, num_class, model.objective_params()),
}
});
Ok(serde_json::to_string_pretty(&value)?)
}
pub fn import_xgboost_json(json: &str) -> Result<BoostedModel> {
let root: Value = serde_json::from_str(json)?;
let learner = field(&root, "learner")?;
let booster = field(learner, "gradient_booster")?;
let booster_name = booster
.get("name")
.and_then(Value::as_str)
.unwrap_or("gbtree");
if booster_name != "gbtree" {
return Err(HessboostError::model_format(format!(
"unsupported gradient_booster `{booster_name}`: only `gbtree` is supported"
)));
}
let model = field(booster, "model")?;
let lmp = field(learner, "learner_model_param")?;
let num_feature = lmp
.get("num_feature")
.and_then(scalar_f64)
.map(|v| v as usize)
.ok_or_else(|| HessboostError::model_format("missing/invalid `num_feature`"))?;
let num_class = lmp
.get("num_class")
.and_then(scalar_f64)
.map_or(0, |v| v as usize);
let objective_json = learner.get("objective");
let objective = objective_json
.and_then(|o| o.get("name"))
.and_then(Value::as_str)
.unwrap_or("reg:squarederror")
.to_string();
let n_outputs = num_class.max(1);
let trees_json = model
.get("trees")
.and_then(Value::as_array)
.ok_or_else(|| HessboostError::model_format("missing `model.trees` array"))?;
let order = round_robin_tree_order(model, trees_json.len(), n_outputs)?;
let mut trees = Vec::with_capacity(trees_json.len());
for &i in &order {
let tree = tree_from_json(&trees_json[i])
.map_err(|e| HessboostError::model_format(format!("tree {i}: {e}")))?;
trees.push(tree);
}
let tree_weights = match model.get("weight_drop") {
None => Vec::new(),
Some(wd) => {
let entries = wd
.as_array()
.ok_or_else(|| HessboostError::model_format("`weight_drop` is not an array"))?;
if entries.len() != trees.len() {
return Err(HessboostError::model_format(format!(
"`weight_drop` has {} entries for {} trees",
entries.len(),
trees.len()
)));
}
let weights = entries
.iter()
.map(|w| scalar_f64(w).map(|w| w as f32))
.collect::<Option<Vec<f32>>>()
.ok_or_else(|| {
HessboostError::model_format("`weight_drop` contains a non-numeric entry")
})?;
order.iter().map(|&i| weights[i]).collect()
}
};
let base_score = lmp
.get("base_score")
.and_then(Value::as_str)
.ok_or_else(|| HessboostError::model_format("missing/invalid `base_score`"))?;
let objective_impl = build_objective(&objective, num_class);
let base_margins = parse_base_score(base_score, objective_impl.as_deref(), n_outputs)?;
let objective_params = objective_params_from_json(&objective, objective_json);
objective_params
.training_params(&objective, num_class)
.build()
.map_err(|e| HessboostError::model_format(format!("invalid objective parameters: {e}")))?;
let imported = BoostedModel::from_parts(
trees,
tree_weights,
base_margins,
ModelSpec {
objective_params,
objective,
num_class,
n_outputs,
n_features: num_feature,
},
);
imported.validate_structure()?;
Ok(imported)
}
fn tree_to_json(id: usize, tree: &RegTree, num_feature: usize) -> Value {
let nodes = tree.nodes();
let n = nodes.len();
let mut left = Vec::with_capacity(n);
let mut right = Vec::with_capacity(n);
let mut split_indices = Vec::with_capacity(n);
let mut split_conditions = Vec::with_capacity(n);
let mut default_left = Vec::with_capacity(n);
let mut base_weights = Vec::with_capacity(n);
let mut loss_changes = Vec::with_capacity(n);
let mut sum_hessian = Vec::with_capacity(n);
let mut split_type = Vec::with_capacity(n);
let mut categories = Vec::<i64>::new();
let mut categories_nodes = Vec::<i64>::new();
let mut categories_segments = Vec::<i64>::new();
let mut categories_sizes = Vec::<i64>::new();
let mut parents = vec![INVALID_NODE; n];
for (i, node) in nodes.iter().enumerate() {
if !node.is_leaf() {
parents[node.left as usize] = i as i32;
parents[node.right as usize] = i as i32;
}
}
for (node_id, node) in nodes.iter().enumerate() {
let categorical = node.is_categorical && !node.is_leaf();
left.push(if categorical { node.right } else { node.left });
right.push(if categorical { node.left } else { node.right });
sum_hessian.push(node.sum_hess);
split_type.push(u32::from(node.is_categorical));
if node.is_categorical {
let cats = &tree.categories()[node.cat_begin as usize..node.cat_end as usize];
categories_nodes.push(node_id as i64);
categories_segments.push(categories.len() as i64);
categories_sizes.push(cats.len() as i64);
categories.extend(cats.iter().map(|&category| i64::from(category)));
}
if node.is_leaf() {
split_indices.push(0u32);
split_conditions.push(node.leaf_value);
base_weights.push(node.leaf_value);
default_left.push(1i32);
loss_changes.push(0.0f32);
} else {
split_indices.push(node.split_feature);
split_conditions.push(node.split_cond);
base_weights.push(0.0f32);
default_left.push(i32::from(if node.is_categorical {
!node.default_left
} else {
node.default_left
}));
loss_changes.push(node.split_gain);
}
}
json!({
"id": id,
"tree_param": {
"num_deleted": "0",
"num_feature": num_feature.to_string(),
"num_nodes": n.to_string(),
"size_leaf_vector": "0",
},
"left_children": left,
"right_children": right,
"parents": parents,
"split_indices": split_indices,
"split_conditions": split_conditions,
"default_left": default_left,
"base_weights": base_weights,
"loss_changes": loss_changes,
"sum_hessian": sum_hessian,
"split_type": split_type,
"categories": categories,
"categories_nodes": categories_nodes,
"categories_segments": categories_segments,
"categories_sizes": categories_sizes,
})
}
fn tree_from_json(tj: &Value) -> Result<RegTree> {
let left = i32_arr(tj, "left_children")
.ok_or_else(|| HessboostError::missing_field("left_children"))?;
let n = left.len();
if n == 0 {
return Err(HessboostError::model_format("tree contains no nodes"));
}
let right = i32_arr(tj, "right_children")
.ok_or_else(|| HessboostError::missing_field("right_children"))?;
if right.len() != n {
return Err(HessboostError::model_format(
"child arrays have different lengths",
));
}
let split_type = strict_nonnegative_integer_array(tj, "split_type")?;
if !split_type.is_empty() && split_type.len() != n {
return Err(HessboostError::model_format(
"`split_type` length does not match the node count",
));
}
let categories = strict_nonnegative_integer_array(tj, "categories")?;
let category_nodes = strict_nonnegative_integer_array(tj, "categories_nodes")?;
let category_segments = strict_nonnegative_integer_array(tj, "categories_segments")?;
let category_sizes = strict_nonnegative_integer_array(tj, "categories_sizes")?;
if category_nodes.len() != category_segments.len()
|| category_nodes.len() != category_sizes.len()
{
return Err(HessboostError::model_format(
"categorical node, segment, and size arrays have different lengths",
));
}
let mut node_categories: Vec<Vec<u32>> = vec![Vec::new(); n];
let mut seen_category_node = vec![false; n];
for slot in 0..category_nodes.len() {
let node = category_nodes[slot] as usize;
let begin = category_segments[slot] as usize;
let size = category_sizes[slot] as usize;
let end = begin
.checked_add(size)
.ok_or_else(|| HessboostError::model_format("categorical segment overflow"))?;
if node >= n || seen_category_node[node] || size == 0 || end > categories.len() {
return Err(HessboostError::model_format("invalid categorical arrays"));
}
seen_category_node[node] = true;
node_categories[node] = categories[begin..end]
.iter()
.map(|&v| {
u32::try_from(v).map_err(|_| HessboostError::model_format("category exceeds u32"))
})
.collect::<Result<Vec<_>>>()?;
}
let split_indices = arr_or_empty(tj, "split_indices");
let split_conditions = arr(tj, "split_conditions", scalar_f64)
.ok_or_else(|| HessboostError::missing_field("split_conditions"))?;
let default_left = arr_or_empty(tj, "default_left");
let base_weights = arr_or_empty(tj, "base_weights");
let sum_hessian = arr_or_empty(tj, "sum_hessian");
let loss_changes = arr_or_empty(tj, "loss_changes");
let at = |v: &[f64], i: usize| v.get(i).copied().unwrap_or(0.0);
let mut nodes = Vec::with_capacity(n);
for i in 0..n {
let sum_hess = at(&sum_hessian, i) as f32;
if left[i] == -1 {
let leaf_value = split_conditions
.get(i)
.copied()
.or_else(|| base_weights.get(i).copied())
.unwrap_or(0.0) as f32;
nodes.push(Node::leaf(leaf_value, sum_hess));
} else {
if left[i] < 0 || right[i] < 0 || left[i] as usize >= n || right[i] as usize >= n {
return Err(HessboostError::model_format(format!(
"node {i} has an invalid child index"
)));
}
let is_categorical = split_type.get(i).copied().unwrap_or(0) != 0;
if is_categorical && !seen_category_node[i] {
return Err(HessboostError::model_format(format!(
"categorical node {i} has no category segment"
)));
}
nodes.push(Node {
split_feature: at(&split_indices, i) as u32,
split_cond: at(&split_conditions, i) as f32,
default_left: if is_categorical {
at(&default_left, i) == 0.0
} else {
at(&default_left, i) != 0.0
},
left: if is_categorical { right[i] } else { left[i] },
right: if is_categorical { left[i] } else { right[i] },
leaf_value: 0.0,
sum_hess,
split_gain: at(&loss_changes, i) as f32,
is_categorical,
cat_begin: 0,
cat_end: 0,
});
}
}
let mut flat_categories: Vec<u32> = Vec::new();
for (i, cats) in node_categories.iter().enumerate() {
if !cats.is_empty() {
nodes[i].cat_begin = flat_categories.len() as u32;
flat_categories.extend(cats);
nodes[i].cat_end = flat_categories.len() as u32;
}
}
let tree: RegTree = serde_json::from_value(json!({
"nodes": nodes,
"categories": flat_categories,
}))?;
Ok(tree)
}
fn build_objective(name: &str, num_class: usize) -> Option<Box<dyn Objective>> {
create_objective(
&ObjectiveParams::defaults_for(name)
.training_params(name, num_class)
.build_unchecked(),
)
.ok()
}
fn format_base_score(margins: &[f32], objective: &dyn Objective) -> String {
let mut stored = margins.to_vec();
if objective.n_outputs() == 1 {
objective.pred_transform(&mut stored);
}
let entries: Vec<String> = stored.iter().map(f32::to_string).collect();
format!("[{}]", entries.join(","))
}
fn parse_base_score(
stored: &str,
objective: Option<&dyn Objective>,
n_outputs: usize,
) -> Result<Vec<f32>> {
let invalid = || HessboostError::model_format(format!("invalid `base_score` `{stored}`"));
let inner = stored
.trim()
.strip_prefix('[')
.and_then(|s| s.strip_suffix(']'))
.ok_or_else(invalid)?;
let values = inner
.split(',')
.map(|v| v.trim().parse::<f32>().ok())
.collect::<Option<Vec<f32>>>()
.ok_or_else(invalid)?;
let values = match values.len() {
1 => vec![values[0]; n_outputs],
len if len == n_outputs => values,
len => {
return Err(HessboostError::model_format(format!(
"`base_score` has {len} entries for {n_outputs} outputs"
)));
}
};
Ok(match objective {
Some(obj) => values.iter().map(|&v| obj.prob_to_margin(v)).collect(),
None => values,
})
}
const SCALE_POS_WEIGHT: (&str, &str) = ("reg_loss_param", "scale_pos_weight");
const MAX_DELTA_STEP: (&str, &str) = ("poisson_regression_param", "max_delta_step");
const TWEEDIE_VARIANCE_POWER: (&str, &str) = ("tweedie_regression_param", "tweedie_variance_power");
const HUBER_SLOPE: (&str, &str) = ("pseudo_huber_param", "huber_slope");
const LAMBDARANK_NUM_PAIR: (&str, &str) = ("lambdarank_param", "lambdarank_num_pair_per_sample");
const SOFTMAX_NUM_CLASS: (&str, &str) = ("softmax_multiclass_param", "num_class");
fn objective_to_json(objective: &str, num_class: usize, params: &ObjectiveParams) -> Value {
let ((block, key), value) = match objective {
"multi:softmax" | "multi:softprob" => (SOFTMAX_NUM_CLASS, num_class.to_string()),
"count:poisson" => (MAX_DELTA_STEP, params.max_delta_step.to_string()),
"reg:tweedie" => (
TWEEDIE_VARIANCE_POWER,
params.tweedie_variance_power.to_string(),
),
"reg:pseudohubererror" => (HUBER_SLOPE, params.huber_slope.to_string()),
"rank:pairwise" | "rank:ndcg" | "rank:map" => (
LAMBDARANK_NUM_PAIR,
params.lambdarank_num_pair_per_sample.to_string(),
),
_ => (SCALE_POS_WEIGHT, params.scale_pos_weight.to_string()),
};
let mut fields = Map::new();
if block == LAMBDARANK_NUM_PAIR.0 {
for (k, v) in [
("lambdarank_bias_norm", "1"),
("lambdarank_normalization", "1"),
("lambdarank_pair_method", "topk"),
("lambdarank_score_normalization", "1"),
("lambdarank_unbiased", "0"),
("ndcg_exp_gain", "1"),
] {
fields.insert(k.to_string(), Value::String(v.to_string()));
}
}
fields.insert(key.to_string(), Value::String(value));
let mut out = Map::with_capacity(2);
out.insert("name".to_string(), Value::String(objective.to_string()));
out.insert(block.to_string(), Value::Object(fields));
Value::Object(out)
}
fn objective_params_from_json(objective: &str, obj: Option<&Value>) -> ObjectiveParams {
let mut params = ObjectiveParams::defaults_for(objective);
let Some(obj) = obj else {
return params;
};
let get =
|(block, key): (&str, &str)| obj.get(block).and_then(|b| b.get(key)).and_then(scalar_f64);
if let Some(v) = get(SCALE_POS_WEIGHT) {
params.scale_pos_weight = v;
}
if let Some(v) = get(MAX_DELTA_STEP) {
params.max_delta_step = v;
}
if let Some(v) = get(TWEEDIE_VARIANCE_POWER) {
params.tweedie_variance_power = v;
}
if let Some(v) = get(HUBER_SLOPE) {
params.huber_slope = v;
}
if let Some(v) = get(LAMBDARANK_NUM_PAIR)
&& v >= 1.0
&& v != f64::from(u32::MAX)
{
params.lambdarank_num_pair_per_sample = v as usize;
}
params
}
fn round_robin_tree_order(model: &Value, n_trees: usize, n_outputs: usize) -> Result<Vec<usize>> {
field(model, "tree_info")?;
let tree_info = strict_nonnegative_integer_array(model, "tree_info")?;
if tree_info.len() != n_trees {
return Err(HessboostError::model_format(format!(
"`tree_info` has {} entries for {n_trees} trees",
tree_info.len()
)));
}
if let Some(bad) = tree_info.iter().find(|&&g| g >= n_outputs as u64) {
return Err(HessboostError::model_format(format!(
"`tree_info` group {bad} out of range for {n_outputs} outputs"
)));
}
let indptr = if model.get("iteration_indptr").is_some() {
let indptr = strict_nonnegative_integer_array(model, "iteration_indptr")?;
let bounded = indptr.first() == Some(&0)
&& indptr.last() == Some(&(n_trees as u64))
&& indptr.windows(2).all(|w| w[0] <= w[1]);
if !bounded {
return Err(HessboostError::model_format(
"`iteration_indptr` must run monotonically from 0 to the number of trees",
));
}
indptr.iter().map(|&i| i as usize).collect::<Vec<_>>()
} else {
let num_parallel_tree = model
.get("gbtree_model_param")
.and_then(|p| p.get("num_parallel_tree"))
.map_or(Some(1.0), scalar_f64)
.filter(|&v| v >= 1.0 && v.fract() == 0.0)
.ok_or_else(|| HessboostError::model_format("invalid `num_parallel_tree`"))?
as usize;
let per_iteration = num_parallel_tree * n_outputs;
if !n_trees.is_multiple_of(per_iteration) {
return Err(HessboostError::model_format(format!(
"{n_trees} trees do not form whole iterations of {per_iteration} \
(num_parallel_tree × outputs)"
)));
}
(0..=n_trees / per_iteration)
.map(|k| k * per_iteration)
.collect()
};
let mut order = Vec::with_capacity(n_trees);
let mut groups: Vec<Vec<usize>> = vec![Vec::new(); n_outputs];
for (iteration, bounds) in indptr.windows(2).enumerate() {
for group in &mut groups {
group.clear();
}
for t in bounds[0]..bounds[1] {
groups[tree_info[t] as usize].push(t);
}
let per_group = groups[0].len();
if groups.iter().any(|g| g.len() != per_group) {
return Err(HessboostError::model_format(format!(
"iteration {iteration}: outputs have unequal tree counts; \
layout is not representable"
)));
}
for round in 0..per_group {
order.extend(groups.iter().map(|g| g[round]));
}
}
Ok(order)
}
fn field<'a>(v: &'a Value, key: &str) -> Result<&'a Value> {
v.get(key).ok_or_else(|| HessboostError::missing_field(key))
}
fn scalar_f64(v: &Value) -> Option<f64> {
match v {
Value::Number(n) => n.as_f64(),
Value::String(s) => s.parse::<f64>().ok(),
Value::Bool(b) => Some(if *b { 1.0 } else { 0.0 }),
_ => None,
}
}
fn arr(v: &Value, key: &str, f: fn(&Value) -> Option<f64>) -> Option<Vec<f64>> {
v.get(key)?
.as_array()
.map(|a| a.iter().map(|e| f(e).unwrap_or(0.0)).collect())
}
fn i32_arr(v: &Value, key: &str) -> Option<Vec<i32>> {
arr(v, key, scalar_f64).map(|a| a.iter().map(|&x| x as i32).collect())
}
fn arr_or_empty(v: &Value, key: &str) -> Vec<f64> {
arr(v, key, scalar_f64).unwrap_or_default()
}
fn strict_nonnegative_integer_array(v: &Value, key: &str) -> Result<Vec<u64>> {
let Some(value) = v.get(key) else {
return Ok(Vec::new());
};
let entries = value
.as_array()
.ok_or_else(|| HessboostError::model_format(format!("`{key}` is not an array")))?;
entries
.iter()
.map(|entry| {
let value = scalar_f64(entry).ok_or_else(|| {
HessboostError::model_format(format!("`{key}` contains a non-numeric entry"))
})?;
if !value.is_finite() || value < 0.0 || value.fract() != 0.0 || value > u64::MAX as f64
{
return Err(HessboostError::model_format(format!(
"`{key}` contains an invalid integer {value}"
)));
}
Ok(value as u64)
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{BoosterKind, TrainingParams};
use crate::data::{DMatrix, FeatureType};
use crate::learner::train;
fn reg_model() -> (BoostedModel, DMatrix) {
let n = 120;
let mut x = Vec::with_capacity(n * 2);
let mut y = Vec::with_capacity(n);
for i in 0..n {
let a = i as f32 / n as f32;
let b = ((i * 7) % n) as f32 / n as f32;
x.push(a);
x.push(b);
y.push(2.0 * a - 3.0 * b + if a > 0.5 { 1.0 } else { -1.0 });
}
let d = DMatrix::from_dense(&x, n, 2)
.unwrap()
.with_labels(&y)
.unwrap();
let params = TrainingParams::builder()
.objective("reg:squarederror")
.max_depth(3)
.eta(0.3)
.build()
.unwrap();
(train(¶ms, &d, 15).unwrap(), d)
}
#[test]
fn roundtrip_reg_preserves_predictions() {
let (model, d) = reg_model();
let before = model.predict(&d).unwrap();
let json = export_xgboost_json(&model).unwrap();
let restored = import_xgboost_json(&json).unwrap();
let after = restored.predict(&d).unwrap();
assert_eq!(restored.num_trees(), model.num_trees());
assert_eq!(restored.n_features(), model.n_features());
assert_eq!(restored.objective(), model.objective());
assert_eq!(before.len(), after.len());
for (a, b) in before.iter().zip(&after) {
assert!((a - b).abs() < 1e-5, "pred drift: {a} vs {b}");
}
}
#[test]
fn roundtrip_binary_preserves_predictions() {
let n = 80;
let x: Vec<f32> = (0..n).map(|i| i as f32 / n as f32).collect();
let y: Vec<f32> = x.iter().map(|&v| f32::from(v > 0.4)).collect();
let d = DMatrix::from_dense(&x, n, 1)
.unwrap()
.with_labels(&y)
.unwrap();
let params = TrainingParams::builder()
.objective("binary:logistic")
.max_depth(3)
.eta(0.3)
.build()
.unwrap();
let model = train(¶ms, &d, 20).unwrap();
let before = model.predict(&d).unwrap();
let json = export_xgboost_json(&model).unwrap();
let restored = import_xgboost_json(&json).unwrap();
assert_eq!(restored.objective(), "binary:logistic");
assert!((restored.base_score() - model.base_score()).abs() < 1e-4);
let after = restored.predict(&d).unwrap();
for (a, b) in before.iter().zip(&after) {
assert!((a - b).abs() < 1e-5, "pred drift: {a} vs {b}");
}
}
fn hand_stump_json() -> &'static str {
r#"{
"version": [3, 4, 1],
"learner": {
"gradient_booster": {
"name": "gbtree",
"model": {
"gbtree_model_param": {"num_parallel_tree": "1", "num_trees": "1"},
"iteration_indptr": [0, 1],
"tree_info": [0],
"trees": [{
"id": 0,
"tree_param": {"num_nodes": "3", "num_feature": "1", "size_leaf_vector": "1"},
"left_children": [1, -1, -1],
"right_children": [2, -1, -1],
"parents": [2147483647, 0, 0],
"split_indices": [0, 0, 0],
"split_conditions": [1.5, 10.0, -10.0],
"default_left": [1, 0, 0],
"base_weights": [0.0, 10.0, -10.0],
"loss_changes": [42.0, 0.0, 0.0],
"sum_hessian": [8.0, 5.0, 3.0],
"split_type": [0, 0, 0]
}]
}
},
"learner_model_param": {
"base_score": "[0E0]", "boost_from_average": "1",
"num_class": "0", "num_feature": "1", "num_target": "1"
},
"objective": {"name": "reg:squarederror", "reg_loss_param": {"scale_pos_weight": "1"}}
}
}"#
}
#[test]
fn import_hand_written_stump_routes_correctly() {
let model = import_xgboost_json(hand_stump_json()).unwrap();
assert_eq!(model.num_trees(), 1);
assert_eq!(model.n_features(), 1);
assert_eq!(model.base_score(), 0.0);
let d = DMatrix::from_dense(&[1.0, 2.0], 2, 1).unwrap();
let margins = model.predict_margin(&d).unwrap();
assert!((margins[0] - 10.0).abs() < 1e-6, "got {}", margins[0]);
assert!((margins[1] + 10.0).abs() < 1e-6, "got {}", margins[1]);
let dm = DMatrix::from_dense(&[f32::NAN], 1, 1).unwrap();
let mm = model.predict_margin(&dm).unwrap();
assert!(
(mm[0] - 10.0).abs() < 1e-6,
"missing routed wrong: {}",
mm[0]
);
}
fn three_class_json(base_score: &str) -> String {
let stump = |id: usize| {
format!(
r#"{{"id": {id}, "tree_param": {{"num_nodes": "1", "num_feature": "2", "size_leaf_vector": "1"}},
"left_children": [-1], "right_children": [-1], "parents": [2147483647],
"split_indices": [0], "split_conditions": [0.0], "default_left": [0],
"base_weights": [0.0], "loss_changes": [0.0], "sum_hessian": [1.0], "split_type": [0]}}"#
)
};
format!(
r#"{{
"version": [3, 4, 1],
"learner": {{
"gradient_booster": {{
"name": "gbtree",
"model": {{
"gbtree_model_param": {{"num_parallel_tree": "1", "num_trees": "3"}},
"tree_info": [0, 1, 2],
"trees": [{}, {}, {}]
}}
}},
"learner_model_param": {{
"base_score": "{base_score}", "boost_from_average": "1",
"num_class": "3", "num_feature": "2", "num_target": "1"
}},
"objective": {{"name": "multi:softprob", "softmax_multiclass_param": {{"num_class": "3"}}}}
}}
}}"#,
stump(0),
stump(1),
stump(2)
)
}
#[test]
fn import_multiclass_vector_intercept_offsets_each_class() {
let model = import_xgboost_json(&three_class_json(
"[5.3293586E-2,-1.3475811E-1,8.146441E-2]",
))
.unwrap();
let expected = [5.329_358_6E-2f32, -1.347_581_1E-1, 8.146_441E-2];
assert_eq!(model.base_scores(), &expected);
let d = DMatrix::from_dense(&[0.0, 0.0, 1.0, 1.0], 2, 2).unwrap();
let margins = model.predict_margin(&d).unwrap();
assert_eq!(margins, [expected, expected].concat());
let uniform = import_xgboost_json(&three_class_json("[5E-1]")).unwrap();
assert_eq!(uniform.base_scores(), &[0.5, 0.5, 0.5]);
}
#[test]
fn malformed_base_score_is_rejected() {
for bad in ["0.5", "[0.1,0.2]", "[a]", "[]", "[0.1,0.2,0.3,0.4]"] {
let err = import_xgboost_json(&three_class_json(bad)).unwrap_err();
assert!(
matches!(err, HessboostError::ModelFormat(_)),
"{bad}: {err}"
);
}
}
#[test]
fn export_writes_xgboost_3_learner_params() {
let (model, _) = reg_model();
let json: Value = serde_json::from_str(&export_xgboost_json(&model).unwrap()).unwrap();
assert_eq!(json["version"], json!([3, 4, 1]));
let lmp = &json["learner"]["learner_model_param"];
assert_eq!(lmp["boost_from_average"], "0");
assert_eq!(lmp["num_target"], "1");
let expected = format!("[{}]", model.base_score());
assert_eq!(lmp["base_score"], expected);
assert_eq!(
json["learner"]["objective"],
json!({"name": "reg:squarederror", "reg_loss_param": {"scale_pos_weight": "1"}})
);
assert!(
json["learner"]["gradient_booster"]["model"]
.get("weight_drop")
.is_none()
);
}
#[test]
fn unsupported_booster_is_rejected() {
let js = r#"{"learner": {"gradient_booster": {"name": "gblinear"},
"learner_model_param": {"num_feature": "3", "base_score": "[0]"}}}"#;
let err = import_xgboost_json(js).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)));
}
#[test]
fn dart_roundtrips_through_weight_drop() {
let (_, d) = reg_model();
let params = TrainingParams::builder()
.booster(BoosterKind::Dart)
.rate_drop(0.5)
.max_depth(3)
.build()
.unwrap();
let model = train(¶ms, &d, 8).unwrap();
assert!(model.has_non_unit_tree_weights());
let before = model.predict(&d).unwrap();
let exported = export_xgboost_json(&model).unwrap();
let json: Value = serde_json::from_str(&exported).unwrap();
assert_eq!(json["learner"]["gradient_booster"]["name"], "gbtree");
let weight_drop = json["learner"]["gradient_booster"]["model"]["weight_drop"]
.as_array()
.unwrap();
assert_eq!(weight_drop.len(), model.num_trees());
let restored = import_xgboost_json(&exported).unwrap();
for t in 0..model.num_trees() {
assert_eq!(restored.tree_weight(t), model.tree_weight(t), "tree {t}");
}
assert_eq!(restored.predict(&d).unwrap(), before);
}
#[test]
fn gblinear_export_is_rejected_and_categorical_roundtrips() {
let (_, d) = reg_model();
let params = TrainingParams::builder()
.booster(BoosterKind::GbLinear)
.build()
.unwrap();
let model = train(¶ms, &d, 3).unwrap();
assert!(export_xgboost_json(&model).is_err());
let categories = [0.0, 1.0, 2.0, 0.0, 1.0, 2.0];
let labels = [1.0, 0.0, 1.0, 1.0, 0.0, 1.0];
let categorical = DMatrix::from_dense(&categories, 6, 1)
.unwrap()
.with_labels(&labels)
.unwrap()
.with_feature_types(&[FeatureType::Categorical])
.unwrap();
let params = TrainingParams::builder().max_depth(2).build().unwrap();
let model = train(¶ms, &categorical, 3).unwrap();
assert!(model.trees().iter().any(|tree| tree.node(0).is_categorical));
let before = model.predict(&categorical).unwrap();
let restored = import_xgboost_json(&export_xgboost_json(&model).unwrap()).unwrap();
assert_eq!(restored.predict(&categorical).unwrap(), before);
}
fn parallel_tree_json(
num_parallel_tree: usize,
tree_info: &[usize],
leaves: &[f32],
iteration_indptr: Option<&[usize]>,
weight_drop: Option<&[f32]>,
) -> String {
let trees: Vec<String> = leaves
.iter()
.enumerate()
.map(|(id, leaf)| {
format!(
r#"{{"id": {id}, "tree_param": {{"num_nodes": "1", "num_feature": "1", "size_leaf_vector": "1"}},
"left_children": [-1], "right_children": [-1], "parents": [2147483647],
"split_indices": [0], "split_conditions": [{leaf}], "default_left": [0],
"base_weights": [{leaf}], "loss_changes": [0.0], "sum_hessian": [1.0], "split_type": [0]}}"#
)
})
.collect();
let extra = |key: &str, values: Option<String>| {
values.map_or(String::new(), |v| format!(r#""{key}": [{v}],"#))
};
let join = |v: &[String]| v.join(", ");
let indptr = extra(
"iteration_indptr",
iteration_indptr.map(|v| join(&v.iter().map(usize::to_string).collect::<Vec<_>>())),
);
let weights = extra(
"weight_drop",
weight_drop.map(|v| join(&v.iter().map(f32::to_string).collect::<Vec<_>>())),
);
format!(
r#"{{
"version": [3, 4, 1],
"learner": {{
"gradient_booster": {{
"name": "gbtree",
"model": {{
"gbtree_model_param": {{"num_parallel_tree": "{num_parallel_tree}", "num_trees": "{}"}},
{indptr}
{weights}
"tree_info": [{}],
"trees": [{}]
}}
}},
"learner_model_param": {{
"base_score": "[0E0]", "boost_from_average": "1",
"num_class": "3", "num_feature": "1", "num_target": "1"
}},
"objective": {{"name": "multi:softprob", "softmax_multiclass_param": {{"num_class": "3"}}}}
}}
}}"#,
leaves.len(),
join(&tree_info.iter().map(usize::to_string).collect::<Vec<_>>()),
join(&trees),
)
}
fn class_margins(json: &str) -> Vec<f32> {
let model = import_xgboost_json(json).unwrap();
let d = DMatrix::from_dense(&[0.0], 1, 1).unwrap();
model.predict_margin(&d).unwrap()
}
#[test]
fn import_regroups_parallel_trees_by_tree_info() {
let leaves = [1.0, 2.0, 10.0, 20.0, 100.0, 200.0];
let tree_info = [0, 0, 1, 1, 2, 2];
let derived = parallel_tree_json(2, &tree_info, &leaves, None, None);
assert_eq!(class_margins(&derived), [3.0, 30.0, 300.0]);
let leaves2 = [leaves, [4.0, 8.0, 40.0, 80.0, 400.0, 800.0]].concat();
let tree_info2 = [tree_info, tree_info].concat();
let explicit = parallel_tree_json(2, &tree_info2, &leaves2, Some(&[0, 6, 12]), None);
assert_eq!(class_margins(&explicit), [15.0, 150.0, 1500.0]);
let weights = [1.0, 0.5, 1.0, 0.5, 1.0, 0.5];
let dart = parallel_tree_json(2, &tree_info, &leaves, None, Some(&weights));
assert_eq!(class_margins(&dart), [2.0, 20.0, 200.0]);
let model = import_xgboost_json(&dart).unwrap();
let restored: Vec<f32> = (0..6).map(|t| model.tree_weight(t)).collect();
assert_eq!(restored, [1.0, 1.0, 1.0, 0.5, 0.5, 0.5]);
let lopsided = parallel_tree_json(2, &[0, 0, 1, 2, 2, 0], &leaves, None, None);
let err = import_xgboost_json(&lopsided).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
let missing = parallel_tree_json(1, &tree_info, &leaves, None, None)
.replace(r#""tree_info": [0, 0, 1, 1, 2, 2],"#, "");
let err = import_xgboost_json(&missing).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
}
fn roundtrip_objective_param(
model: &BoostedModel,
block: &str,
key: &str,
expected: &str,
) -> BoostedModel {
let exported = export_xgboost_json(model).unwrap();
let json: Value = serde_json::from_str(&exported).unwrap();
assert_eq!(
json["learner"]["objective"][block][key], expected,
"{block}.{key}"
);
import_xgboost_json(&exported).unwrap()
}
#[test]
fn objective_params_roundtrip_through_parameter_blocks() {
let n = 40;
let x: Vec<f32> = (0..n).map(|i| i as f32 / n as f32).collect();
let counts: Vec<f32> = (0..n).map(|i| (i % 4) as f32).collect();
let data = |labels: &[f32]| {
DMatrix::from_dense(&x, n, 1)
.unwrap()
.with_labels(labels)
.unwrap()
};
let fit = |builder: crate::config::TrainingParamsBuilder, d: &DMatrix| {
train(&builder.max_depth(2).build().unwrap(), d, 2).unwrap()
};
let d = data(&counts);
let tweedie = fit(
TrainingParams::builder()
.objective("reg:tweedie")
.tweedie_variance_power(1.2),
&d,
);
let back = roundtrip_objective_param(
&tweedie,
"tweedie_regression_param",
"tweedie_variance_power",
"1.2",
);
assert_eq!(back.objective_params().tweedie_variance_power, 1.2);
let poisson = fit(
TrainingParams::builder()
.objective("count:poisson")
.max_delta_step(0.3),
&d,
);
let back = roundtrip_objective_param(
&poisson,
"poisson_regression_param",
"max_delta_step",
"0.3",
);
assert_eq!(back.objective_params().max_delta_step, 0.3);
let huber = fit(
TrainingParams::builder()
.objective("reg:pseudohubererror")
.huber_slope(2.5),
&d,
);
let back = roundtrip_objective_param(&huber, "pseudo_huber_param", "huber_slope", "2.5");
assert_eq!(back.objective_params().huber_slope, 2.5);
let binary: Vec<f32> = x.iter().map(|&v| f32::from(v > 0.6)).collect();
let logistic = fit(
TrainingParams::builder()
.objective("binary:logistic")
.scale_pos_weight(3.0),
&data(&binary),
);
let back = roundtrip_objective_param(&logistic, "reg_loss_param", "scale_pos_weight", "3");
assert_eq!(back.objective_params().scale_pos_weight, 3.0);
let ranked = data(&counts).with_group_sizes(&[20, 20]).unwrap();
let ranker = fit(
TrainingParams::builder()
.objective("rank:ndcg")
.lambdarank_num_pair_per_sample(5),
&ranked,
);
let back = roundtrip_objective_param(
&ranker,
"lambdarank_param",
"lambdarank_num_pair_per_sample",
"5",
);
assert_eq!(back.objective_params().lambdarank_num_pair_per_sample, 5);
let exported = export_xgboost_json(&ranker).unwrap().replace(
r#""lambdarank_num_pair_per_sample": "5""#,
r#""lambdarank_num_pair_per_sample": "4294967295""#,
);
let unset = import_xgboost_json(&exported).unwrap();
assert_eq!(unset.objective_params().lambdarank_num_pair_per_sample, 32);
}
#[test]
fn custom_objective_export_is_rejected() {
use crate::learner::train_with_objective;
use crate::objective::{CustomObjective, GradPair};
let (_, d) = reg_model();
let params = TrainingParams::builder()
.objective("custom:test")
.max_depth(2)
.build()
.unwrap();
let obj = CustomObjective::new("custom:test", 1, 0.0, "rmse", |preds, labels, w, out| {
for i in 0..preds.len() {
let wi = w.map_or(1.0, |ws| ws[i]);
out[i] = GradPair::new((preds[i] - labels[i]) * wi, wi);
}
});
let model = train_with_objective(¶ms, &d, 2, &obj).unwrap();
let err = export_xgboost_json(&model).unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
}
}