Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,8 @@ validated `WalkerConfig`. The existing panic-based loaders remain available.
The supported subset includes scalar regression, ranking and binary
classification, `identity`/`sigmoid`, sum/average aggregation, and scalar base
scores. Binary Treelite v4 is recommended; native XGBoost JSON must first be
converted through Treelite. Limits are 64 features and 128 rows per group.
converted through Treelite. Inputs have up to 64 features; groups can have any
number of rows.
See [Treelite loading](docs/treelite-loading.md) for export examples, the exact
output and precision policy, errors, limits, and the caller's grouping contract.

Expand Down Expand Up @@ -201,9 +202,13 @@ only on the released CSVs.

## Correctness

LightGBM (f64) predictions match treelite GTIL within 1e-14; XGBoost (f32)
predictions match native XGBoost within 1e-5 (`benchmarks/tests/correctness.rs`).
Partial evaluation and the full walk agree within 1e-15.
Predictions are the correctly rounded sums of the leaf values: leaves are added as
exact fixed-point integers and rounded once, so tree order, layout and ablation
modes cannot change a prediction. LightGBM (f64) predictions match treelite GTIL
within 1e-13; XGBoost (f32) predictions match native XGBoost within 1e-5
(`benchmarks/tests/correctness.rs`). GTIL and the full walk add leaves in f64 in tree
order, and their rounding grows with the partial sums: on a 1,008-step survival
panel, where margins reach ±500, GTIL is up to 1.1e-14 from the correctly rounded sum.

## Citation

Expand Down
29 changes: 26 additions & 3 deletions benchmarks/src/external.rs
Original file line number Diff line number Diff line change
Expand Up @@ -220,6 +220,11 @@ pub struct LightGBMBench {
}

impl LightGBMBench {
/// Output of the most recent `predict_group` call.
pub fn last_output(&self, n: usize) -> &[f64] {
&self.result_buf[..n]
}

/// Load LightGBM from a shared library and a model file.
///
/// `max_group_width` controls the result buffer size (one prediction per row).
Expand Down Expand Up @@ -298,7 +303,7 @@ impl ExternalMethod for LightGBMBench {
let data_ptr = unsafe { data.as_ptr().add(start * n_cols) };
let mut out_len: i64 = 0;
unsafe {
(self.predict_fn)(
let rc = (self.predict_fn)(
self.handle,
data_ptr.cast::<c_void>(),
1, // C_API_DTYPE_FLOAT64
Expand All @@ -312,6 +317,7 @@ impl ExternalMethod for LightGBMBench {
&raw mut out_len,
self.result_buf.as_mut_ptr(),
);
assert_eq!(rc, 0, "LGBM_BoosterPredictForMat failed (rc={rc})");
}
}

Expand Down Expand Up @@ -358,6 +364,18 @@ pub struct XGBoostBench {
f32_buf: Vec<f32>,
// Pre-allocated JSON config string (reused across calls).
predict_config: CString,
last_out: *const f32,
last_n: usize,
}

impl XGBoostBench {
/// Output of the most recent `predict_group` call (valid until the next call).
pub fn last_output(&self) -> Vec<f32> {
if self.last_out.is_null() {
return Vec::new();
}
unsafe { std::slice::from_raw_parts(self.last_out, self.last_n).to_vec() }
}
}

impl XGBoostBench {
Expand Down Expand Up @@ -444,7 +462,7 @@ impl XGBoostBench {
let f32_buf = vec![0.0f32; 128 * n_cols];
// Prediction config: normal prediction, no iteration limit.
let predict_config = CString::new(
r#"{"type":0,"training":false,"iteration_range":[0,0],"strict_shape":false}"#,
r#"{"type":0,"training":false,"iteration_begin":0,"iteration_end":0,"strict_shape":false,"missing":NaN}"#,
)
.unwrap();

Expand All @@ -457,6 +475,8 @@ impl XGBoostBench {
free_booster_fn,
f32_buf,
predict_config,
last_out: std::ptr::null(),
last_n: 0,
})
}
}
Expand Down Expand Up @@ -492,7 +512,7 @@ impl ExternalMethod for XGBoostBench {
let mut out_shape: *const u64 = std::ptr::null();
let mut out_dim: u64 = 0;
let mut out_result: *const f32 = std::ptr::null();
(self.predict_from_dense_fn)(
let rc = (self.predict_from_dense_fn)(
self.handle,
array_cstr.as_ptr(),
self.predict_config.as_ptr(),
Expand All @@ -501,6 +521,9 @@ impl ExternalMethod for XGBoostBench {
&raw mut out_dim,
&raw mut out_result,
);
assert_eq!(rc, 0, "XGBoosterPredictFromDense failed (rc={rc})");
self.last_out = out_result;
self.last_n = nrow;
}
}

Expand Down
142 changes: 115 additions & 27 deletions benchmarks/tests/correctness.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,15 +8,19 @@ use treewalker_gbdt::forest::{Forest, ThresholdType};
// Tolerances
// ---------------------------------------------------------------------------

/// F64 models (LightGBM): GTIL reference is exact. Only FP accumulation noise.
const TOL_F64: f64 = 1e-14;
/// F64 models (LightGBM). GTIL adds leaves in f64 in tree order, so its rounding
/// grows with the partial sums: on the 1,008-step FLCHAIN panel, where margins reach
/// ±500, it is up to 1.1e-14 from the correctly rounded sum that `predict` returns.
const TOL_F64: f64 = 1e-13;

/// F32 models (XGBoost): f32 leaf accumulation noise between TW (f64 sum) and
/// native XGBoost (f32 sum). Split decisions are identical (f32 comparison).
const TOL_F32: f64 = 1e-5;

/// Partial eval vs full walk: must be bitwise identical (same code path, same precision).
const TOL_EXACT: f64 = 1e-15;
/// Partial eval vs full walk. `predict` returns the correctly rounded sum of the leaf
/// values; the full walk adds them in f64 in tree order, so the two differ by the full
/// walk's rounding: the same accumulation noise as the GTIL comparison.
const TOL_FULL_WALK: f64 = TOL_F64;

/// Current importer bounds JSON memory; larger models use streaming binary.
const SIMD_JSON_MAX_BYTES: u64 = 64 * 1024 * 1024;
Expand Down Expand Up @@ -223,6 +227,42 @@ fn max_diff(a: &[f64], b: &[f64]) -> f64 {
.fold(0.0f64, f64::max)
}

/// Whether every group's declared monotonic features are monotonic (ignoring NaN).
/// The tiled chunked-G cells repeat a 16-step panel inside each group, so they break
/// this caller contract on purpose.
fn honors_monotonic_contract(
config: &WalkerConfig,
data: &[f64],
n_rows: usize,
offsets: Option<&[usize]>,
) -> bool {
let nf = config.n_features;
let mut ok = true;
for_each_group(config.max_group_width, n_rows, offsets, |s, e| {
for f in (0..nf).filter(|&f| config.is_mono_inc(f) || config.is_mono_dec(f)) {
let values: Vec<f64> = (s..e)
.map(|r| data[r * nf + f])
.filter(|v| !v.is_nan())
.collect();
ok &= values.windows(2).all(|w| {
if config.is_mono_inc(f) {
w[0] <= w[1]
} else {
w[0] >= w[1]
}
});
}
});
ok
}

fn assert_same_bits(actual: &[f64], expected: &[f64], label: &str) {
assert_eq!(actual.len(), expected.len(), "{label}: length");
if let Some(r) = (0..actual.len()).find(|&r| actual[r].to_bits() != expected[r].to_bits()) {
panic!("{label}: row {r}: {} != {}", actual[r], expected[r]);
}
}

// ---------------------------------------------------------------------------
// Core correctness — runs on EVERY artifact
// ---------------------------------------------------------------------------
Expand Down Expand Up @@ -267,7 +307,7 @@ fn test_partial_matches_full() {

let d = max_diff(&partial, &full);
eprintln!("{}: partial_vs_full={d:.2e}", cfg.label);
assert!(d < TOL_EXACT, "{} partial vs full: {d:.2e}", cfg.label);
assert!(d < TOL_FULL_WALK, "{} partial vs full: {d:.2e}", cfg.label);
}
}

Expand Down Expand Up @@ -314,7 +354,7 @@ fn assert_partial_matches_full_on_obs(forest: &mut Forest, data: &[f64], obs: us

let d = max_diff(&full[start..end], &partial[start..end]);
assert!(
d < TOL_EXACT,
d < TOL_FULL_WALK,
"Partial vs full max_diff={d:.2e} on observation {obs}"
);
}
Expand Down Expand Up @@ -442,17 +482,40 @@ fn test_ablation_all_configs() {
),
];

// Every mode reaches the same leaves with the same rows, so predictions are
// bitwise identical to the default.
for cfg in &all_configs() {
let (data, n_rows) = load_test_data(&cfg.param_dir);
let offsets = cfg.group_offsets.as_deref();
let reference = predict_all(
&mut load_forest(&cfg.param_dir, cfg.framework),
&data,
n_rows,
offsets,
);
let monotonic = {
let forest = load_forest(&cfg.param_dir, cfg.framework);
honors_monotonic_contract(&forest.config, &data, n_rows, offsets)
};
for &(mode, ablation) in ablations {
// Without precompute, monotonic features are partitioned by prefix/suffix
// scans that are only correct when the data honor the declared contract.
let scans = ablation.disable_varying_precompute && !ablation.disable_monotonic;
if scans && !monotonic {
eprintln!(
"{}/{mode}: skipped, data break the monotonic contract",
cfg.label
);
continue;
}
let mut forest = load_forest(&cfg.param_dir, cfg.framework);
forest.config.ablation = ablation;
let (data, n_rows) = load_test_data(&cfg.param_dir);

let partial = predict_all(&mut forest, &data, n_rows, cfg.group_offsets.as_deref());
let full = predict_full_all(&forest, &data, n_rows, cfg.group_offsets.as_deref());

let d = max_diff(&partial, &full);
assert!(d < TOL_EXACT, "{}/{mode} max_diff={d:.2e}", cfg.label);
let mut ablated = vec![0.0f64; n_rows];
let width = forest.config.max_group_width;
for_each_group(width, n_rows, offsets, |s, e| {
forest.predict_with_stats(&data, &mut ablated, s, e);
});
assert_same_bits(&ablated, &reference, &format!("{}/{mode}", cfg.label));
}
eprintln!("{}: all ablations ok", cfg.label);
}
Expand Down Expand Up @@ -498,18 +561,27 @@ fn test_parse_configs() {
// disable_predicate_dedup tested separately — large models exceed u16 limit.
];

// With exact sums, tree order and node layout cannot change a prediction.
for cfg in &all_configs() {
let (data, n_rows) = load_test_data(&cfg.param_dir);
let offsets = cfg.group_offsets.as_deref();
let mut reference_forest = load_forest(&cfg.param_dir, cfg.framework);
let reference = predict_all(&mut reference_forest, &data, n_rows, offsets);
for &(mode, ref pc) in parse_configs {
let mut forest = Forest::load_with_config(
model_file(&cfg.param_dir, cfg.framework),
cfg.param_dir.join("walker_config.json"),
pc,
);
let (data, n_rows) = load_test_data(&cfg.param_dir);
let partial = predict_all(&mut forest, &data, n_rows, cfg.group_offsets.as_deref());
let full = predict_full_all(&forest, &data, n_rows, cfg.group_offsets.as_deref());
let d = max_diff(&partial, &full);
assert!(d < TOL_EXACT, "{}/{mode} max_diff={d:.2e}", cfg.label);
let partial = predict_all(&mut forest, &data, n_rows, offsets);
let label = format!("{}/{mode}", cfg.label);
if reference_forest.exact_sums() {
assert_same_bits(&partial, &reference, &label);
} else {
let full = predict_full_all(&forest, &data, n_rows, offsets);
let d = max_diff(&partial, &full);
assert!(d < TOL_FULL_WALK, "{label} max_diff={d:.2e}");
}
}
eprintln!("{}: all parse configs ok", cfg.label);
}
Expand Down Expand Up @@ -626,7 +698,7 @@ fn test_cat_out_of_range_partial_matches_full() {
let d = max_diff(&full_results[start..end], &partial_results[start..end]);
eprintln!("{}: cat_out_of_range partial_vs_full={d:.2e}", cfg.label);
assert!(
d < TOL_EXACT,
d < TOL_FULL_WALK,
"{}: cat out-of-range partial vs full max_diff={d:.2e}",
cfg.label
);
Expand Down Expand Up @@ -679,7 +751,7 @@ fn test_width_32_boundary() {
nf,
);
assert!(
d < TOL_EXACT,
d < TOL_FULL_WALK,
"Width-32 boundary: partial vs full max_diff={d:.2e} — u32 mask arithmetic may be wrong"
);

Expand All @@ -688,7 +760,7 @@ fn test_width_32_boundary() {
forest.predict(&data, &mut ws_results, 0, 32);
let d2 = max_diff(&full_results, &ws_results);
assert!(
d2 < TOL_EXACT,
d2 < TOL_FULL_WALK,
"Width-32 boundary (workspace): max_diff={d2:.2e}"
);
}
Expand Down Expand Up @@ -859,7 +931,7 @@ fn test_wide_group_partial_matches_full(target_width: usize) {
let total_rows = n_synthetic_groups * target_width;
assert_eq!(synth_data.len(), total_rows * nf);

// Predict with partial eval (exercises u64/u128 mask path).
// Predict with partial eval (exercises the mask width chosen for target_width).
let mut results_partial = vec![0.0f64; total_rows];
for g in 0..n_synthetic_groups {
let start = g * target_width;
Expand All @@ -881,27 +953,43 @@ fn test_wide_group_partial_matches_full(target_width: usize) {
first_cfg.framework,
);
assert!(
diff < TOL_EXACT,
diff < TOL_FULL_WALK,
"wide_group(width={target_width}): partial eval diverges from full walk: max_diff={diff:.2e}"
);
}

#[test]
fn test_wide_group_48_u64() {
fn test_wide_group_48() {
test_wide_group_partial_matches_full(48);
}

#[test]
fn test_wide_group_64_u64() {
fn test_wide_group_64() {
test_wide_group_partial_matches_full(64);
}

#[test]
fn test_wide_group_96_u128() {
fn test_wide_group_96() {
test_wide_group_partial_matches_full(96);
}

#[test]
fn test_wide_group_128_u128() {
fn test_wide_group_128() {
test_wide_group_partial_matches_full(128);
}

#[test]
fn test_wide_group_200() {
test_wide_group_partial_matches_full(200);
}

#[test]
fn test_wide_group_1000() {
test_wide_group_partial_matches_full(1000);
}

/// Wider than one 1,024-row piece.
#[test]
fn test_wide_group_2500() {
test_wide_group_partial_matches_full(2500);
}
15 changes: 12 additions & 3 deletions docs/treelite-loading.md
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,7 @@ should construct a `Forest`.
| Types | Matching float64 thresholds/leaves, or matching float32 thresholds/leaves |
| Numeric operators | float64 `<` and `<=`; float32 `<` |
| Categories | Nonnegative category IDs up to 8159; membership in either child direction |
| Features / grouped rows | 1–64 features; configured maximum 1–128 rows |
| Features / grouped rows | 1–64 features; any positive maximum group width |

Multiclass/multiple targets, vector leaves (including singleton vectors), other
postprocessors/operators/type combinations, nonfinite leaves/base scores,
Expand All @@ -100,13 +100,22 @@ output = 1 / (1 + exp(-alpha * margin)) # sigmoid
The base score is stored on the forest, rather than distributed across leaves.
This can change the last few rounding bits from older TreeWalker versions.

`predict` and `predict_with_stats` compute `tree_sum` exactly: every leaf value is
an integer multiple of 2^-e for one scale e per model, the scaled leaves add as
128-bit integers, and the sum is rounded to float64 once. The result is the float64
nearest to the true sum and does not depend on tree order, node layout or group
width. This needs the scaled sums to fit in 126 bits; a model with a non-finite leaf
or leaf exponents too far apart adds in float64 in tree order instead, which
`Forest::exact_sums` reports. `predict_full` always adds in float64 in tree order,
so it can differ from `predict` in the last bits.

Float64 numerical inputs compare in float64. `<` thresholds are normalized to
`<= next_down(threshold)`, including signed zero and positive infinity. Float64
`< -infinity` and all NaN thresholds are explicitly rejected because this
normalization cannot represent them exactly. Float64 `<= -infinity` is valid.
Float32 models round numerical inputs and thresholds to float32 before `<`;
categorical inputs are also rounded to float32. Leaves are promoted and summed
in float64, so native XGBoost bitwise parity is not promised. Use float32 input
categorical inputs are also rounded to float32. Leaves are promoted to float64
and summed as above, so native XGBoost bitwise parity is not promised. Use float32 input
to Treelite GTIL when testing this policy; GTIL with float64 input can select a
different branch near a float32 threshold.

Expand Down
Loading