hessboost 0.2.1

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
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
//! Regression tree representation and prediction.
//!
//! A tree is a flat array of [`Node`]s (node `0` is the root). Internal nodes
//! carry a numeric split `x[feature] < threshold`. Instances whose feature is
//! *missing* follow the node's `default_left` direction, implementing XGBoost's
//! sparsity-aware routing. Leaf nodes carry the raw leaf weight (the learning
//! rate is applied by the boosting loop, not baked into the tree).
//!
//! A *vector-leaf* tree (`multi_strategy = multi_output_tree`, XGBoost's
//! `MultiTargetTree`) shares one split structure across `K > 1` outputs and
//! stores a weight vector per leaf ([`RegTree::leaf_vector`]); its scalar
//! [`Node::leaf_value`]s are unused (zero).

use crate::data::DMatrix;
use crate::error::HessboostError;
use crate::tree::linear::{LinearLeaves, UncheckedLinearLeaves};
use crate::tree::{SplitTest, split_goes_left};
use serde::{Deserialize, Serialize};

/// Sentinel used in child pointers to mark "no child" (i.e. a leaf).
const NO_CHILD: i32 = -1;

/// A single tree node.
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Node {
    /// Split feature index (meaningful only for internal nodes).
    pub split_feature: u32,
    /// Split threshold: an instance goes left when `value < split_cond`.
    pub split_cond: f32,
    /// Direction taken by instances with a missing value at this node.
    pub default_left: bool,
    /// Left child index, or `-1` for a leaf.
    pub left: i32,
    /// Right child index, or `-1` for a leaf.
    pub right: i32,
    /// Leaf weight (used only for leaves).
    pub leaf_value: f32,
    /// Sum of Hessians routed through this node (for cover-based importance/SHAP).
    pub sum_hess: f32,
    /// Loss reduction (gain) achieved by this node's split (0 for leaves).
    pub split_gain: f32,
    /// Whether this internal node splits on a categorical feature by set
    /// membership rather than a numeric threshold. `false` for numeric splits
    /// and leaves.
    pub is_categorical: bool,
    /// For a categorical node, the start index into the owning tree's category
    /// list of the categories routed left. Unused (`0`) otherwise.
    pub cat_begin: u32,
    /// For a categorical node, the end index (exclusive) into the owning tree's
    /// category list of the categories routed left. Unused (`0`) otherwise.
    pub cat_end: u32,
}

impl Node {
    /// A fresh leaf node with the given weight and cover.
    pub(crate) fn leaf(value: f32, sum_hess: f32) -> Self {
        Node {
            split_feature: 0,
            split_cond: 0.0,
            default_left: true,
            left: NO_CHILD,
            right: NO_CHILD,
            leaf_value: value,
            sum_hess,
            split_gain: 0.0,
            is_categorical: false,
            cat_begin: 0,
            cat_end: 0,
        }
    }

    /// Whether this node is a leaf.
    #[inline]
    pub fn is_leaf(&self) -> bool {
        self.left == NO_CHILD
    }
}

/// How an internal node routes rows: the feature it tests, the test, and the
/// direction rows missing that feature take. Passed to [`RegTree::expand`].
#[derive(Debug, Clone, Copy)]
pub(crate) struct SplitRule<'a> {
    feature: u32,
    test: SplitTest<'a>,
    default_left: bool,
}

impl<'a> SplitRule<'a> {
    /// A numeric split: `x[feature] < threshold` goes left, missing values
    /// follow `default_left`.
    pub(crate) fn numeric(feature: u32, threshold: f32, default_left: bool) -> Self {
        SplitRule {
            feature,
            test: SplitTest::Threshold(threshold),
            default_left,
        }
    }

    /// A categorical (set-membership) split: instances whose value of
    /// `feature` is one of `cats_left` go left, other present categories go
    /// right, and missing values follow `default_left`.
    pub(crate) fn categorical(feature: u32, cats_left: &'a [u32], default_left: bool) -> Self {
        SplitRule {
            feature,
            test: SplitTest::Categories(cats_left),
            default_left,
        }
    }
}

/// Initial weight and cover (Hessian sum) of a leaf created by
/// [`RegTree::expand`].
#[derive(Debug, Clone, Copy)]
pub(crate) struct ChildLeaf {
    value: f32,
    sum_hess: f32,
}

impl ChildLeaf {
    pub(crate) fn new(value: f32, sum_hess: f32) -> Self {
        ChildLeaf { value, sum_hess }
    }
}

/// A regression tree: a flat node array with node `0` as the root.
///
/// Serde deserialization validates the tree's own structure (child links in
/// range, every node reached exactly once from the root, finite values,
/// category ranges inside the category pool, leaf vectors and linear leaves
/// consistent with the nodes) and refuses a malformed tree, so every
/// deserialized tree can be traversed. Feature indices are checked against a
/// feature count only by the owning [`BoostedModel`](crate::model::BoostedModel).
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "UncheckedRegTree")]
pub struct RegTree {
    nodes: Vec<Node>,
    /// Flat pool of category values routed left by categorical nodes. Node
    /// `n` (when `n.is_categorical`) owns `categories[n.cat_begin..n.cat_end]`.
    /// Empty for trees with no categorical splits.
    categories: Vec<u32>,
    /// Outputs per leaf of a vector-leaf tree (XGBoost's `size_leaf_vector`,
    /// `> 1`), or `0` for a scalar tree.
    size_leaf_vector: usize,
    /// Vector-leaf weights laid out `[node][output]` (`size_leaf_vector` per
    /// node; internal nodes hold zeros). Empty for scalar trees.
    leaf_vectors: Vec<f32>,
    /// Per-leaf linear models of a `linear_tree` tree ([`LinearLeaves`]);
    /// `None` for constant-leaf trees (every tree unless `linear_tree` is on).
    linear: Option<LinearLeaves>,
}

/// The serialized fields of a [`RegTree`] (same names and layout), before
/// validation. `RegTree`'s `Deserialize` goes through it; the native JSON
/// reader keeps it unchecked so the model validates each tree against its
/// feature count ([`UncheckedRegTree::into_unchecked`]). A scalar tree may
/// omit `size_leaf_vector` (`0`) and `leaf_vectors` (empty); a multi-output
/// model requires `size_leaf_vector`
/// ([`UncheckedRegTree::states_leaf_width`]). `linear` is required (`null`
/// for constant leaves): nothing else marks a linear-leaf tree, so an
/// omitted one would silently predict its constant leaf values.
#[derive(Deserialize)]
pub(crate) struct UncheckedRegTree {
    nodes: Vec<Node>,
    categories: Vec<u32>,
    size_leaf_vector: Option<usize>,
    #[serde(default)]
    leaf_vectors: Vec<f32>,
    // Present but nullable: a plain `Option` field would default when absent.
    #[serde(deserialize_with = "Option::deserialize")]
    linear: Option<UncheckedLinearLeaves>,
}

impl UncheckedRegTree {
    /// Whether the tree stored `size_leaf_vector` rather than defaulting it.
    pub(crate) fn states_leaf_width(&self) -> bool {
        self.size_leaf_vector.is_some()
    }

    /// The tree as stored, unvalidated: the caller validates it
    /// ([`RegTree::is_valid_for_features`]).
    pub(crate) fn into_unchecked(self) -> RegTree {
        RegTree::from_parts(
            self.nodes,
            self.categories,
            self.size_leaf_vector.unwrap_or(0),
            self.leaf_vectors,
            self.linear.map(UncheckedLinearLeaves::into_unchecked),
        )
    }
}

impl TryFrom<UncheckedRegTree> for RegTree {
    type Error = HessboostError;

    fn try_from(unchecked: UncheckedRegTree) -> Result<Self, Self::Error> {
        let tree = unchecked.into_unchecked();
        // No feature count bounds a standalone tree: traversal reads features
        // through the caller's accessor, which answers any index.
        if tree.is_valid_for_features(usize::MAX) {
            Ok(tree)
        } else {
            Err(HessboostError::model_format("tree contains invalid nodes"))
        }
    }
}

impl RegTree {
    /// Create an empty tree with a placeholder root leaf (node `0`), ready to
    /// be grown by a builder.
    pub(crate) fn with_root(sum_hess: f32) -> Self {
        Self::from_scalar_parts(vec![Node::leaf(0.0, sum_hess)], Vec::new())
    }

    /// Assemble a constant-leaf scalar tree from its node array and the flat
    /// category pool its categorical nodes index. The caller validates the
    /// result ([`RegTree::is_valid_for_features`]).
    pub(crate) fn from_scalar_parts(nodes: Vec<Node>, categories: Vec<u32>) -> Self {
        RegTree {
            nodes,
            categories,
            size_leaf_vector: 0,
            leaf_vectors: Vec::new(),
            linear: None,
        }
    }

    /// Assemble a tree from every stored part, as the native binary format
    /// keeps them: `size_leaf_vector` is `0` for a scalar tree, and
    /// `leaf_vectors` holds `size_leaf_vector` weights per node. The caller
    /// validates the result ([`RegTree::is_valid_for_features`]).
    pub(crate) fn from_parts(
        nodes: Vec<Node>,
        categories: Vec<u32>,
        size_leaf_vector: usize,
        leaf_vectors: Vec<f32>,
        linear: Option<LinearLeaves>,
    ) -> Self {
        RegTree {
            nodes,
            categories,
            size_leaf_vector,
            leaf_vectors,
            linear,
        }
    }

    /// The stored leaf-vector width (`0` for a scalar tree) and weights, as
    /// [`RegTree::from_parts`] takes them.
    pub(crate) fn leaf_vector_parts(&self) -> (usize, &[f32]) {
        (self.size_leaf_vector, &self.leaf_vectors)
    }

    /// Create a vector-leaf tree with `n_outputs > 1` weights per leaf and a
    /// placeholder (all-zero) root leaf.
    pub(crate) fn with_vector_root(n_outputs: usize, sum_hess: f32) -> Self {
        debug_assert!(n_outputs > 1);
        RegTree {
            nodes: vec![Node::leaf(0.0, sum_hess)],
            categories: Vec::new(),
            size_leaf_vector: n_outputs,
            leaf_vectors: vec![0.0; n_outputs],
            linear: None,
        }
    }

    /// Outputs per leaf: `1` for a scalar tree, `K > 1` for a vector-leaf
    /// tree.
    #[inline]
    pub fn size_leaf_vector(&self) -> usize {
        self.size_leaf_vector.max(1)
    }

    /// Whether this is a vector-leaf (multi-output) tree.
    #[inline]
    pub fn is_vector_leaf(&self) -> bool {
        self.size_leaf_vector > 1
    }

    /// The weights of leaf `nid`, one per output: the leaf's vector for a
    /// vector-leaf tree, the single [`Node::leaf_value`] otherwise.
    ///
    /// With a leaf id from [`leaf_id_with`](Self::leaf_id_with) or
    /// [`leaf_id_dense`](Self::leaf_id_dense) this is the tree's output for
    /// a row (as XGBoost's `GetLeafIndex` and `LeafValue`), except in a
    /// linear-leaf tree, where the leaf predicts its
    /// [`linear_leaves`](Self::linear_leaves) model and the weight is only its
    /// fallback for rows missing one of the model's features.
    #[inline]
    pub fn leaf_vector(&self, nid: usize) -> &[f32] {
        if self.is_vector_leaf() {
            let k = self.size_leaf_vector;
            &self.leaf_vectors[nid * k..(nid + 1) * k]
        } else {
            std::slice::from_ref(&self.nodes[nid].leaf_value)
        }
    }

    /// Set the weight vector of vector-leaf tree node `nid`.
    pub(crate) fn set_leaf_vector(&mut self, nid: usize, values: &[f32]) {
        let k = self.size_leaf_vector;
        debug_assert!(k > 1 && values.len() == k);
        self.leaf_vectors[nid * k..(nid + 1) * k].copy_from_slice(values);
    }

    /// Give two freshly pushed child nodes their (zero) leaf vectors.
    fn grow_leaf_vectors(&mut self) {
        if self.is_vector_leaf() {
            self.leaf_vectors
                .resize(self.nodes.len() * self.size_leaf_vector, 0.0);
        }
    }

    /// Number of nodes (internal + leaf).
    #[inline]
    pub fn num_nodes(&self) -> usize {
        self.nodes.len()
    }

    /// Number of leaf nodes.
    pub fn num_leaves(&self) -> usize {
        self.nodes.iter().filter(|n| n.is_leaf()).count()
    }

    pub(crate) fn is_valid_for_features(&self, n_features: usize) -> bool {
        let locally_valid = !self.nodes.is_empty()
            && self.nodes.iter().all(|node| {
                node.sum_hess.is_finite()
                    && node.leaf_value.is_finite()
                    && node.split_cond.is_finite()
                    && node.split_gain.is_finite()
                    // A leaf has no right child either: the XGBoost
                    // importer refuses any other entry, so an accepted model
                    // must export to what it imports.
                    && ((node.is_leaf() && node.right == NO_CHILD)
                        || ((node.split_feature as usize) < n_features
                            && node.left >= 0
                            && node.right >= 0
                            && (node.left as usize) < self.nodes.len()
                            && (node.right as usize) < self.nodes.len()))
                    // Scalar trees carry the category set of categorical
                    // leaves too (XGBoost keeps a pruned split's set), and
                    // XGBoost refuses an empty one.
                    && (!node.is_categorical
                        || (node.is_leaf() && self.is_vector_leaf())
                        || (node.cat_begin < node.cat_end
                            && (node.cat_end as usize) <= self.categories.len()))
            })
            && self
                .linear
                .as_ref()
                .is_none_or(|linear| linear.is_valid(&self.nodes, n_features));
        if !locally_valid {
            return false;
        }
        // `size_leaf_vector` is untrusted: an overflowing product is invalid.
        if self.is_vector_leaf()
            && (self.nodes.len().checked_mul(self.size_leaf_vector)
                != Some(self.leaf_vectors.len())
                || self.leaf_vectors.iter().any(|w| !w.is_finite()))
        {
            return false;
        }
        // A leaf is a scalar constant, a vector, or a scalar linear model:
        // vector-leaf consumers ignore linear payloads, so both at once
        // would make them disagree with `predict_row`.
        if self.size_leaf_vector == 1
            || (!self.is_vector_leaf() && !self.leaf_vectors.is_empty())
            || (self.is_vector_leaf() && self.linear.is_some())
        {
            return false;
        }
        let mut seen = vec![false; self.nodes.len()];
        let mut stack = vec![0usize];
        while let Some(node_id) = stack.pop() {
            if seen[node_id] {
                return false;
            }
            seen[node_id] = true;
            let node = &self.nodes[node_id];
            if !node.is_leaf() {
                stack.push(node.left as usize);
                stack.push(node.right as usize);
            }
        }
        seen.into_iter().all(|visited| visited)
    }

    /// Read-only access to the node array.
    #[inline]
    pub fn nodes(&self) -> &[Node] {
        &self.nodes
    }

    /// Flat pool of categories routed left by categorical nodes.
    #[inline]
    pub(crate) fn categories(&self) -> &[u32] {
        &self.categories
    }

    /// The categories categorical node `node` of this tree routes left.
    #[inline]
    pub(crate) fn node_categories(&self, node: &Node) -> &[u32] {
        &self.categories[node.cat_begin as usize..node.cat_end as usize]
    }

    /// The per-leaf linear models, when this is a linear-leaf tree (trained
    /// with `linear_tree`). Such a leaf predicts its linear model, or its
    /// constant `leaf_value` for rows missing one of the model's features.
    #[inline]
    pub fn linear_leaves(&self) -> Option<&LinearLeaves> {
        self.linear.as_ref()
    }

    /// Attach fitted leaf linear models.
    pub(crate) fn set_linear_leaves(&mut self, linear: LinearLeaves) {
        self.linear = Some(linear);
    }

    /// Access a node by id.
    #[inline]
    pub fn node(&self, id: usize) -> &Node {
        &self.nodes[id]
    }

    /// Turn leaf `nid` into an internal node routing rows by `split`, with
    /// two new child leaves `left` and `right`. Returns `(left_id, right_id)`.
    ///
    /// Both builders overwrite the child values in their finalize pass (which
    /// recomputes every leaf from stored stats and bounds), so the values here
    /// are placeholders on that path — but the parameters stay: directly built
    /// trees (tests, learners) rely on them as the real leaf weights.
    pub(crate) fn expand(
        &mut self,
        nid: usize,
        split: SplitRule<'_>,
        left: ChildLeaf,
        right: ChildLeaf,
    ) -> (usize, usize) {
        match split.test {
            SplitTest::Threshold(cond) => self.nodes[nid].split_cond = cond,
            SplitTest::Categories(cats_left) => {
                let begin = self.categories.len() as u32;
                self.categories.extend_from_slice(cats_left);
                let n = &mut self.nodes[nid];
                n.is_categorical = true;
                n.cat_begin = begin;
                n.cat_end = self.categories.len() as u32;
            }
        }
        let left_id = self.nodes.len();
        let right_id = left_id + 1;
        let n = &mut self.nodes[nid];
        n.split_feature = split.feature;
        n.default_left = split.default_left;
        n.left = left_id as i32;
        n.right = right_id as i32;
        self.nodes.push(Node::leaf(left.value, left.sum_hess));
        self.nodes.push(Node::leaf(right.value, right.sum_hess));
        self.grow_leaf_vectors();
        (left_id, right_id)
    }

    /// Set a leaf's weight (used to finalize leaf values after growth).
    pub(crate) fn set_leaf_value(&mut self, nid: usize, value: f32) {
        self.nodes[nid].leaf_value = value;
    }

    /// Record the loss reduction achieved by an internal node's split.
    pub(crate) fn set_split_gain(&mut self, nid: usize, gain: f32) {
        self.nodes[nid].split_gain = gain;
    }

    /// Record the Hessian sum (cover) of the instances reaching node `nid`.
    pub(crate) fn set_sum_hess(&mut self, nid: usize, sum_hess: f32) {
        self.nodes[nid].sum_hess = sum_hess;
    }

    /// Multiply every leaf weight by `factor`. Used to apply the learning rate
    /// (shrinkage) so that stored trees already carry their scaled contribution,
    /// matching XGBoost's saved-model semantics. Leaf linear models are scaled
    /// with them.
    pub(crate) fn scale_leaves(&mut self, factor: f32) {
        let k = self.size_leaf_vector;
        for (id, n) in self.nodes.iter_mut().enumerate() {
            if n.is_leaf() {
                n.leaf_value *= factor;
                if k > 1 {
                    for w in &mut self.leaf_vectors[id * k..(id + 1) * k] {
                        *w *= factor;
                    }
                }
            }
        }
        if let Some(linear) = &mut self.linear {
            linear.scale(f64::from(factor));
        }
    }

    /// Add `delta` to every (scalar) leaf weight: EBM's centering of a tree
    /// by its mean over the training rows.
    pub(crate) fn shift_leaves(&mut self, delta: f32) {
        for n in &mut self.nodes {
            if n.is_leaf() {
                n.leaf_value += delta;
            }
        }
    }

    /// Route a single feature vector (via an accessor) to its leaf id.
    ///
    /// `get` returns `None` for a missing feature. Generic over the accessor so
    /// the same code serves dense rows, sparse rows, and SHAP traversals. Each
    /// level loads its node once (re-indexing through `child` measured
    /// ~7% slower).
    pub fn leaf_id_with(&self, get: impl Fn(u32) -> Option<f32>) -> usize {
        let nodes = &self.nodes[..];
        let mut nid = 0usize;
        loop {
            let node = &nodes[nid];
            if node.is_leaf() {
                return nid;
            }
            nid = if self.goes_left(node, get(node.split_feature)) {
                node.left as usize
            } else {
                node.right as usize
            };
        }
    }

    /// The child of internal node `nid` that an instance whose split-feature
    /// value is `value` (`None` = missing) descends to.
    #[inline]
    pub(crate) fn child(&self, nid: usize, value: Option<f32>) -> usize {
        let node = &self.nodes[nid];
        if self.goes_left(node, value) {
            node.left as usize
        } else {
            node.right as usize
        }
    }

    /// Whether an instance whose split-feature value is `value` (`None` =
    /// missing) goes left at internal node `node` of this tree.
    #[inline]
    pub(crate) fn goes_left(&self, node: &Node, value: Option<f32>) -> bool {
        // Categories are integer-coded; membership in the left set routes
        // left, everything else (present, not in set) right.
        let test = if node.is_categorical {
            SplitTest::Categories(self.node_categories(node))
        } else {
            SplitTest::Threshold(node.split_cond)
        };
        split_goes_left(value, node.default_left, test)
    }

    /// Route a dense feature row (indexed by feature id, `missing` sentinel for
    /// absent values) to its leaf id.
    #[inline]
    pub fn leaf_id_dense(&self, row: &[f32], missing: f32) -> usize {
        self.leaf_id_with(|f| {
            let v = row[f as usize];
            (!crate::data::is_missing(v, missing)).then_some(v)
        })
    }

    /// The raw output of row `row` of `data` in this *scalar* tree: its
    /// leaf's weight, or the leaf's linear model for linear-leaf trees.
    /// Crate-internal: vector-leaf trees hold one weight per output, which
    /// callers read with [`RegTree::leaf_vector`].
    pub(crate) fn predict_row(&self, data: &DMatrix, row: usize) -> f32 {
        debug_assert!(!self.is_vector_leaf(), "predict_row on a vector-leaf tree");
        let get = |f: u32| data.get(row, f as usize);
        let leaf = self.leaf_id_with(get);
        let constant = self.nodes[leaf].leaf_value;
        match &self.linear {
            Some(linear) => linear.predict(leaf, constant, get),
            None => constant,
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    /// Build a small tree by hand:
    /// root: feature 0 < 0.5 ? left : right, missing -> left
    ///   left leaf = -1.0, right leaf = +2.0
    fn stump() -> RegTree {
        let mut t = RegTree::with_root(10.0);
        t.expand(
            0,
            SplitRule::numeric(0, 0.5, true),
            ChildLeaf::new(-1.0, 5.0),
            ChildLeaf::new(2.0, 5.0),
        );
        t
    }

    /// A deserialized tree is traversable: serde refuses child links out of
    /// range, cycles, and unreachable nodes (which traversal would follow
    /// out of bounds or forever) and a leaf vector of the wrong length,
    /// including a width whose storage size overflows `usize`.
    #[test]
    fn deserialization_refuses_malformed_trees() {
        let doc = serde_json::to_value(stump()).unwrap();
        assert_eq!(
            serde_json::from_value::<RegTree>(doc.clone()).unwrap(),
            stump()
        );
        for (node, field, value) in [(0, "left", 3), (0, "right", 0), (0, "left", 2)] {
            let mut bad = doc.clone();
            bad["nodes"][node][field] = value.into();
            assert!(
                serde_json::from_value::<RegTree>(bad).is_err(),
                "{field} = {value}"
            );
        }
        let mut vector = serde_json::to_value(RegTree::with_vector_root(2, 1.0)).unwrap();
        assert!(serde_json::from_value::<RegTree>(vector.clone()).is_ok());
        vector["leaf_vectors"] = serde_json::json!([0.0]);
        assert!(serde_json::from_value::<RegTree>(vector).is_err());
        // Three nodes of `usize::MAX` weights each overflow the count.
        let mut wide = doc;
        wide["size_leaf_vector"] = usize::MAX.into();
        wide["leaf_vectors"] = serde_json::json!([]);
        assert!(serde_json::from_value::<RegTree>(wide).is_err());
    }

    #[test]
    fn routing_numeric() {
        let t = stump();
        // value 0.2 < 0.5 -> left leaf -1.0
        assert_eq!(t.leaf_id_with(|_| Some(0.2)), 1);
        // value 0.9 >= 0.5 -> right leaf +2.0
        assert_eq!(t.leaf_id_with(|_| Some(0.9)), 2);
    }

    #[test]
    fn routing_missing_follows_default() {
        let t = stump();
        // missing -> default_left = true -> left leaf
        assert_eq!(t.leaf_id_with(|_| None), 1);
        assert_eq!(t.node(1).leaf_value, -1.0);
    }

    #[test]
    fn routing_categorical_set_membership() {
        // Categorical split: categories {0, 2} go left, everything else right.
        let mut t = RegTree::with_root(10.0);
        t.expand(
            0,
            SplitRule::categorical(0, &[0, 2], false),
            ChildLeaf::new(-1.0, 5.0),
            ChildLeaf::new(2.0, 5.0),
        );
        assert!(t.node(0).is_categorical);
        // In-set categories route left (leaf 1, value -1.0).
        assert_eq!(t.leaf_id_with(|_| Some(0.0)), 1);
        assert_eq!(t.leaf_id_with(|_| Some(2.0)), 1);
        // Out-of-set present categories route right (leaf 2, value 2.0).
        assert_eq!(t.leaf_id_with(|_| Some(1.0)), 2);
        assert_eq!(t.leaf_id_with(|_| Some(3.0)), 2);
        // Unseen category also routes right (not in the left set).
        assert_eq!(t.leaf_id_with(|_| Some(9.0)), 2);
        // Missing follows default_left = false -> right.
        assert_eq!(t.leaf_id_with(|_| None), 2);
    }

    #[test]
    fn predict_row_dense() {
        let t = stump();
        let d = DMatrix::from_dense(&[0.1, 0.9], 2, 1).unwrap();
        assert_eq!(t.predict_row(&d, 0), -1.0);
        assert_eq!(t.predict_row(&d, 1), 2.0);
        assert_eq!(t.num_leaves(), 2);
        assert_eq!(t.num_nodes(), 3);
    }

    /// A vector-leaf tree carrying a linear-leaf payload is invalid: vector
    /// consumers would drop the linear models `predict_row` uses.
    #[test]
    fn vector_leaves_refuse_linear_payload() {
        let linear: LinearLeaves = serde_json::from_str(
            r#"{"offsets":[0,1],"intercepts":[0.5],"features":[0],"coeffs":[2.0]}"#,
        )
        .unwrap();
        let mut scalar = RegTree::with_root(1.0);
        scalar.set_linear_leaves(linear.clone());
        assert!(scalar.is_valid_for_features(1));
        let mut vector = RegTree::with_vector_root(2, 1.0);
        assert!(vector.is_valid_for_features(1));
        vector.set_linear_leaves(linear);
        assert!(!vector.is_valid_for_features(1));
    }

    /// A leaf (no left child) with a right child is invalid: the XGBoost
    /// formats refuse it, so accepting it would load a model that cannot be
    /// exported and imported again.
    #[test]
    fn a_leaf_refuses_a_right_child() {
        let mut tree = RegTree::with_root(1.0);
        assert!(tree.is_valid_for_features(1));
        tree.nodes[0].right = -2;
        assert!(!tree.is_valid_for_features(1));
    }
}