hessboost 0.2.0

Fast, deterministic gradient boosting (GBDT) in Rust: conformal intervals, explainable boosting machines, distributional boosting, tree-based diffusion, and XGBoost model interchange
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
//! One boosting round: gradients, row subsets, and the trees of an iteration
//! (or `process_type=update`'s refresh of them).

use super::dart::{dart_new_tree_weight, finish_dart, round_gradients};
use super::margins::{MarginCaches, TreeOutput};
use super::prepare::{Prepared, TrainContext, TreeSample, approx_index};
use super::row_sampling::{gradient_sampling, iteration_row_subsets, make_column_sampler};
use crate::config::{Device, Refresh, TrainingParams};
use crate::data::ghist::GHistIndex;
use crate::error::Result;
use crate::model::BoostedModel;
use crate::objective::GradPair;
use crate::rng::Rng;
use crate::training::multi_output;
use crate::training::refresh::refresh_tree;
use crate::training::sampling::{GradientSample, gradient_based_sample};
use crate::training::sglb::LeafRenewal;
use crate::tree::RegTree;
use crate::tree::builder::LeafRows;
use crate::tree::reuse::ReuseSet;
use crate::tree::sampler::ColumnSampler;
use rayon::prelude::*;
use std::sync::OnceLock;

/// What the tree-growing and refresh rounds update: the ensemble, its
/// margin caches, the gradient buffers, and the reuse dictionary.
pub(super) struct RoundState<'a> {
    pub(super) model: BoostedModel,
    pub(super) margins: MarginCaches<'a>,
    /// Every output's gradients, `[row][n_out]`.
    pub(super) gpair: Vec<GradPair>,
    /// One output's gradients gathered from `gpair` (empty for
    /// single-output objectives, which read `gpair` directly).
    pub(super) gpair_k: Vec<GradPair>,
    pub(super) reuse: Option<ReuseSet>,
    /// SGLB's noisy structure gradients, `[row][n_out]` (empty otherwise).
    pub(super) noisy_gpair: Vec<GradPair>,
    /// Every training row, ascending: the rows of an unsampled round.
    pub(super) all_rows: Vec<u32>,
}

/// `process_type=update`: refresh iteration `iteration`'s trees of `queue`
/// output by output with `refresh`'s options, from the gradients of the
/// already refreshed ones, and re-append them.
pub(super) fn refresh_round(
    run: &TrainContext,
    queue: &mut [RegTree],
    refresh: Refresh,
    iteration: usize,
    state: &mut RoundState,
) -> Result<()> {
    let TrainContext {
        params,
        dtrain,
        info,
        objective,
        ..
    } = *run;
    let n_out = objective.n_outputs();
    let parallel = params.num_parallel_tree;
    // Gradients from the already refreshed iterations; iteration `i`'s trees
    // are then refreshed in place, output by output.
    objective.gradient_info_at(&state.margins.train, info, &mut state.gpair, iteration);
    multi_output::reject_split_gradient(objective, iteration, &state.gpair)?;
    let per_iteration = n_out * parallel;
    for slot in 0..per_iteration {
        let k = slot / parallel;
        let gk = gather_output(&state.gpair, &mut state.gpair_k, n_out, k);
        let mut tree = std::mem::replace(
            &mut queue[iteration * per_iteration + slot],
            RegTree::with_root(0.0),
        );
        refresh_tree(&mut tree, dtrain, gk, params, refresh, tree_eta(params));
        state.margins.add_tree(&tree, TreeOutput::Scalar(k), None);
        state.model.push_tree_weighted(tree, 1.0);
    }
    Ok(())
}

/// Grow one iteration's scalar-leaf trees (gbtree or DART) with `prepared`
/// and append them.
pub(super) fn grow_round(
    run: &TrainContext,
    prepared: &Prepared,
    iteration: usize,
    state: &mut RoundState,
) -> Result<()> {
    let TrainContext {
        params,
        dtrain,
        objective,
        ..
    } = *run;
    let n = dtrain.n_rows();
    let n_out = objective.n_outputs();
    let parallel = params.num_parallel_tree;
    // 1. Gradients from the current margins (all outputs at once), DART's
    //    from the ensemble minus this round's dropout set.
    let (mut rng, dropped) = round_gradients(
        run,
        &state.model,
        iteration,
        &state.margins.train,
        &mut state.gpair,
    );
    multi_output::reject_split_gradient(objective, iteration, &state.gpair)?;
    let weight = dart_new_tree_weight(dropped.as_ref(), params);
    // SGLB: the structure is searched on noisy gradients; `state.gpair`
    // keeps the noise-free ones the leaves are re-estimated from.
    let structure: &[GradPair] = match run.langevin {
        Some(langevin) => {
            langevin.structure_gradients(&state.gpair, iteration, &mut state.noisy_gpair)
        }
        None => &state.gpair,
    };

    // 2. Row subsets (uniform, class-balanced, or by query), drawn before
    //    the trees and shared across the per-output fits.
    let row_subsets = iteration_row_subsets(
        params,
        prepared.samples_per_forest(),
        run.rows,
        &state.all_rows,
        &mut rng,
    );
    // An output's gradient-based sample, when its whole forest shares one.
    let mut forest_sample = None;
    let forest_indices = prepared.forest_indices(n_out, parallel);
    let grow = GrowRound {
        run,
        prepared,
        gpair: structure,
        clean_gpair: &state.gpair,
        n_out,
        iteration,
        forest_indices: &forest_indices,
    };

    // 3. `num_parallel_tree` trees per output from the same gradients,
    //    output-major like XGBoost's layout.
    let slots: Vec<TreeSlot> = (0..n_out * parallel)
        .map(|slot| {
            let row_subset = row_subsets.rows(slot % parallel);
            // Retaining the final row partitions replaces per-row tree
            // traversals of the raw feature matrix with one sequential pass
            // per leaf: the training margin update's, when every row took
            // part, and the linear-leaf fit's. Linear leaves use them only
            // when they equal raw routing (see `rows_route_like_trees`).
            let (routed, linear_rows) = match prepared {
                Prepared::Hist {
                    rows_route_like_trees,
                    ..
                } => (true, params.linear_tree.is_some() && *rows_route_like_trees),
                // The exact builder routes rows by their raw values, as
                // prediction does.
                Prepared::Exact(_) => (true, false),
                Prepared::Approx { .. } => (false, false),
            };
            let margin_rows = routed
                && dropped.is_none()
                && row_subset.len() == n
                && !gradient_sampling(params)
                && (params.linear_tree.is_none() || linear_rows);
            TreeSlot {
                output: slot / parallel,
                parallel: slot % parallel,
                rows: row_subset,
                // Leaf re-estimation (SGLB) reads the partitions too.
                capture_rows: margin_rows || linear_rows || (routed && run.langevin.is_some()),
                margin_rows,
            }
        })
        .collect();
    // The trees of an iteration share the round's gradients and do not read
    // each other. Without gradient-based sampling or a reuse dictionary, each
    // tree's RNG draws are its column sampler and rounding seed, drawn here
    // in slot order as the sequential path draws them; the trees are then
    // grown in parallel. A GPU backend stages one tree's gradients at a
    // time, so it keeps the sequential path.
    let trees: Vec<(RegTree, Vec<LeafRows>)> = if slots.len() > 1
        && state.reuse.is_none()
        && !gradient_sampling(params)
        && params.device == Device::Cpu
        && rayon::current_num_threads() > 1
    {
        // The first tree's cuts, before any tree reads them.
        prepared.fill_approx_cache(run, gather_output(structure, &mut state.gpair_k, n_out, 0));
        let draws: Vec<(ColumnSampler, u64)> = slots
            .iter()
            .map(|_| {
                let sampler = make_column_sampler(dtrain, params, &mut rng);
                (sampler, quantization_seed(params, &mut rng))
            })
            .collect();
        // Every output's gradients gathered once, output-major, for all of
        // its parallel trees (single-output objectives read `gpair`).
        let gathered = gather_outputs(structure, n_out);
        let output_gpair = |k: usize| {
            if n_out == 1 {
                structure
            } else {
                &gathered[k * n..(k + 1) * n]
            }
        };
        // Build the forests' shared `approx` indices here, before the
        // parallel trees read them: an index built inside a tree task would
        // run its own parallel loops while other tasks wait on it, and a
        // worker waiting there can steal a task that waits on the same
        // index again.
        for (k, index) in forest_indices.iter().enumerate() {
            index.get_or_init(|| approx_index(params, dtrain, output_gpair(k), false));
        }
        slots
            .par_iter()
            .zip(draws)
            .map(|(slot, (mut sampler, rounding_seed))| {
                let gk = output_gpair(slot.output);
                let sample = TreeSample {
                    gpair: gk,
                    rows: slot.rows,
                    forest_index: grow.forest_index(slot.output),
                };
                grow_sampled_tree(&grow, slot, sample, &mut sampler, rounding_seed, None)
            })
            .collect()
    } else {
        slots
            .iter()
            .map(|slot| {
                fit_output_tree(
                    &grow,
                    slot,
                    &mut state.gpair_k,
                    &mut rng,
                    &mut forest_sample,
                    state.reuse.as_mut(),
                )
            })
            .collect::<Result<_>>()?
    };

    for (slot, (tree, leaf_rows)) in slots.iter().zip(trees) {
        // A dropout round's gradients come from the ensemble, not the margin
        // caches, which `finish_dart` recomputes.
        if dropped.is_none() {
            // The builder's final row partitions already identify the training
            // leaves when every row took part in growing the tree.
            let captured = slot.margin_rows.then_some(leaf_rows.as_slice());
            state
                .margins
                .add_tree(&tree, TreeOutput::Scalar(slot.output), captured);
        }
        state.model.push_tree_weighted(tree, weight);
    }
    if let Some(dropped) = &dropped {
        finish_dart(&mut state.model, params, dropped, &mut state.margins);
    }
    Ok(())
}

/// What the boosting rounds do to the ensemble.
pub(super) enum RoundPlan {
    /// Grow new trees with the prepared builder state.
    Grow(Prepared),
    /// `process_type=update`: refresh the queued trees of the initial model,
    /// one iteration per round, with the refresh updater's options.
    Refresh(Vec<RegTree>, Refresh),
}

/// The learning rate applied to each new tree: `eta / num_parallel_tree`
/// (XGBoost divides the rate across a forest so a whole iteration moves by
/// `eta`), in `f32` as XGBoost's `learning_rate` is.
pub(super) fn tree_eta(params: &TrainingParams) -> f32 {
    params.eta as f32 / params.num_parallel_tree as f32
}

/// Borrow the gradient slice for output `k`: the whole buffer for
/// single-output objectives, otherwise gather output `k`'s pairs into
/// `scratch` (length `n`) and borrow that.
pub(super) fn gather_output<'a>(
    gpair: &'a [GradPair],
    scratch: &'a mut [GradPair],
    n_out: usize,
    k: usize,
) -> &'a [GradPair] {
    if n_out == 1 {
        gpair
    } else {
        for (r, dst) in scratch.iter_mut().enumerate() {
            *dst = gpair[r * n_out + k];
        }
        scratch
    }
}

/// Every output's gradients of `gpair` (`[row][n_out]`) gathered
/// output-major (`[output][row]`, as [`gather_output`] gathers one); empty
/// for a single output.
fn gather_outputs(gpair: &[GradPair], n_out: usize) -> Vec<GradPair> {
    if n_out == 1 {
        return Vec::new();
    }
    let n = gpair.len() / n_out;
    let mut out = vec![GradPair::default(); gpair.len()];
    out.par_chunks_exact_mut(n)
        .enumerate()
        .for_each(|(k, column)| {
            for (r, dst) in column.iter_mut().enumerate() {
                *dst = gpair[r * n_out + k];
            }
        });
    out
}

/// One tree-growing boosting iteration: what each of its trees reads.
struct GrowRound<'a> {
    run: &'a TrainContext<'a>,
    prepared: &'a Prepared,
    /// Every output's gradients the structures are searched on,
    /// `[row][n_out]` (with Langevin noise under SGLB).
    gpair: &'a [GradPair],
    /// The noise-free gradients SGLB re-estimates the leaves from (the
    /// same as `gpair` otherwise).
    clean_gpair: &'a [GradPair],
    n_out: usize,
    /// The model's absolute iteration index.
    iteration: usize,
    /// One gradient index per output, shared by that output's forest
    /// ([`Prepared::forest_indices`]); empty when trees build their own.
    forest_indices: &'a [OnceLock<GHistIndex>],
}

impl GrowRound<'_> {
    /// The gradient index output `output`'s forest shares, if any.
    fn forest_index(&self, output: usize) -> Option<&OnceLock<GHistIndex>> {
        self.forest_indices.get(output)
    }
}

/// Which tree of an iteration to grow: parallel tree `parallel` of output
/// `output`, on the uniform row subset `rows`, keeping its leaves' rows when
/// `capture_rows` (see [`Prepared::build_tree`]) and adding it to the
/// training margins from them when `margin_rows`.
struct TreeSlot<'a> {
    output: usize,
    parallel: usize,
    rows: &'a [u32],
    capture_rows: bool,
    margin_rows: bool,
}

/// Fit the tree `slot` of the iteration `grow`: gather that output's
/// gradient slice (into `scratch` for multi-output objectives), apply
/// gradient-based row sampling when configured (per tree, as XGBoost's hist
/// updater does, or once per output forest under `approx`, kept in
/// `forest_sample` by the forest's first tree for the rest), derive its
/// column sampler, build the tree, fit linear leaves when configured, and
/// shrink its leaves by `eta / num_parallel_tree`. The caller owns the round
/// RNG (already seeded and salted), the reuse dictionary, and what happens
/// to the tree (margin updates, contribution weight).
fn fit_output_tree(
    grow: &GrowRound,
    slot: &TreeSlot,
    scratch: &mut [GradPair],
    rng: &mut Rng,
    forest_sample: &mut Option<GradientSample>,
    reuse: Option<&mut ReuseSet>,
) -> Result<(RegTree, Vec<LeafRows>)> {
    let TrainContext { params, dtrain, .. } = *grow.run;
    let (prepared, n_out) = (grow.prepared, grow.n_out);
    let gk: &[GradPair] = gather_output(grow.gpair, scratch, n_out, slot.output);
    let own;
    let sampled = if !gradient_sampling(params) {
        None
    } else if prepared.samples_per_forest() {
        if slot.parallel == 0 {
            *forest_sample = gradient_based_sample(gk, 1, params.subsample, rng)?;
        }
        forest_sample.as_ref()
    } else {
        own = gradient_based_sample(gk, 1, params.subsample, rng)?;
        own.as_ref()
    };
    let (gk, rows) = match sampled {
        Some(s) => (s.gpair.as_slice(), s.rows.as_slice()),
        None => (gk, slot.rows),
    };
    let mut sampler = make_column_sampler(dtrain, params, rng);
    let rounding_seed = quantization_seed(params, rng);
    let sample = TreeSample {
        gpair: gk,
        rows,
        forest_index: grow.forest_index(slot.output),
    };
    Ok(grow_sampled_tree(
        grow,
        slot,
        sample,
        &mut sampler,
        rounding_seed,
        reuse,
    ))
}

/// The part of [`fit_output_tree`] after its RNG draws: build the tree on
/// `sample`, fit linear leaves when configured, and shrink its leaves.
fn grow_sampled_tree(
    grow: &GrowRound,
    slot: &TreeSlot,
    sample: TreeSample,
    sampler: &mut ColumnSampler,
    rounding_seed: u64,
    reuse: Option<&mut ReuseSet>,
) -> (RegTree, Vec<LeafRows>) {
    let TrainContext { params, dtrain, .. } = *grow.run;
    let TreeSample {
        gpair: gk, rows, ..
    } = sample;
    let (mut tree, leaf_rows) = grow.prepared.build_tree(
        grow.run,
        sample,
        sampler,
        reuse,
        rounding_seed,
        slot.capture_rows,
    );
    if let Some(langevin) = grow.run.langevin {
        let at = LeafRenewal {
            data: dtrain,
            gpair: grow.clean_gpair,
            n_out: grow.n_out,
            rows,
            leaf_rows: &leaf_rows,
            iteration: grow.iteration,
            tree: slot.output * params.num_parallel_tree + slot.parallel,
        };
        langevin.renew_leaves(&mut tree, TreeOutput::Scalar(slot.output), &at);
    }
    // LightGBM keeps the first iteration's trees constant.
    if let Some(linear_tree) = params.linear_tree
        && grow.iteration > 0
    {
        let lambda = linear_tree.lambda();
        if leaf_rows.is_empty() {
            crate::tree::linear_fit::fit_linear_leaves(&mut tree, dtrain, gk, rows, lambda);
        } else {
            crate::tree::linear_fit::fit_captured_linear_leaves(
                &mut tree, dtrain, gk, &leaf_rows, lambda,
            );
        }
    }
    tree.scale_leaves(tree_eta(params));
    (tree, leaf_rows)
}

/// The stochastic-rounding seed of one quantized tree, drawn from the
/// iteration's RNG after the tree's column sampler, so every tree of an
/// iteration (outputs and parallel trees alike) rounds independently and
/// continued training resumes the same streams. Draws nothing unless
/// `use_quantized_grad` is on, leaving the default RNG streams untouched.
fn quantization_seed(params: &TrainingParams, rng: &mut Rng) -> u64 {
    if params.quantized.is_some() {
        rng.next_u64()
    } else {
        0
    }
}