use super::compact::LANES;
const MAX_DEPTH: usize = 16;
const MAX_SLOTS_PER_LEAF: usize = 4;
pub(crate) enum ArenaNode {
Leaf(f32),
Numeric {
slot: u32,
key: u32,
first: u32,
},
Other,
}
#[derive(Debug, Clone)]
pub(crate) struct SymmetricTree {
levels: Vec<(u32, u32)>,
leaves: Vec<u32>,
values: Vec<f32>,
}
impl SymmetricTree {
fn detect(root: u32, node: impl Fn(u32) -> ArenaNode) -> Option<Self> {
let mut levels: Vec<(u32, u32)> = Vec::new();
let mut n_leaves = 0usize;
let mut frontier = vec![root];
let mut next = Vec::new();
while !frontier.is_empty() {
next.clear();
let mut level = None;
for &id in &frontier {
match node(id) {
ArenaNode::Leaf(_) => n_leaves += 1,
ArenaNode::Numeric { slot, key, first } => {
if *level.get_or_insert((slot, key)) != (slot, key) {
return None;
}
next.extend([first, first + 1]);
}
ArenaNode::Other => return None,
}
}
if let Some(level) = level {
if levels.len() == MAX_DEPTH {
return None;
}
levels.push(level);
}
std::mem::swap(&mut frontier, &mut next);
}
let depth = levels.len();
if depth < 2 || (1usize << depth) > MAX_SLOTS_PER_LEAF * n_leaves {
return None;
}
let leaves: Vec<u32> = (0..1u32 << depth)
.map(|pattern| {
let mut id = root;
for d in 0..depth {
match node(id) {
ArenaNode::Numeric { first, .. } => {
id = first + ((pattern >> (depth - 1 - d)) & 1);
}
_ => break,
}
}
id
})
.collect();
let values = leaves
.iter()
.map(|&id| match node(id) {
ArenaNode::Leaf(value) => value,
_ => unreachable!("patterns end at leaves"),
})
.collect();
Some(SymmetricTree {
levels,
leaves,
values,
})
}
#[inline]
pub(crate) fn walk(
&self,
lanes: &[u32],
groups: usize,
group_len: usize,
sink: &mut impl FnMut(usize, u32),
) {
self.for_each_pattern(lanes, groups, group_len, |row, p| {
sink(row, self.leaves[p as usize]);
});
}
#[inline]
pub(crate) fn accumulate(
&self,
lanes: &[u32],
groups: usize,
group_len: usize,
weight: f32,
out: &mut [f32],
stride: usize,
) {
self.for_each_pattern(lanes, groups, group_len, |row, p| {
out[row * stride] += weight * self.values[p as usize];
});
}
#[inline(always)]
fn for_each_pattern(
&self,
lanes: &[u32],
groups: usize,
group_len: usize,
mut sink: impl FnMut(usize, u32),
) {
for (g, grp) in lanes.chunks_exact(group_len).take(groups).enumerate() {
for (j, &p) in self.patterns(grp).iter().enumerate() {
sink(g * LANES + j, p);
}
}
}
#[inline(always)]
fn patterns(&self, grp: &[u32]) -> [u32; LANES] {
let mut pattern = [0u32; LANES];
for &(slot, key) in &self.levels {
let keys: &[u32; LANES] = grp[slot as usize..]
.first_chunk()
.expect("split slot addresses a full lane run");
for (p, &k) in pattern.iter_mut().zip(keys) {
*p = (*p << 1) | u32::from(k > key);
}
}
pattern
}
}
#[derive(Debug, Clone, Default)]
pub(crate) struct SymmetricTables {
trees: Vec<Option<SymmetricTree>>,
}
impl SymmetricTables {
pub(crate) fn push(&mut self, root: u32, node: impl Fn(u32) -> ArenaNode) {
self.trees.push(SymmetricTree::detect(root, node));
}
#[inline]
pub(crate) fn get(&self, t: usize) -> Option<&SymmetricTree> {
self.trees[t].as_ref()
}
}
#[cfg(test)]
mod tests {
use crate::config::{GrowPolicy, TrainingParams, TreeMethod};
use crate::model::Iterations;
use crate::test_support::labeled_dense;
use crate::training::train;
use crate::tree::compact::{CompactForest, LANES, LaneBlock, split_lanes};
use crate::tree::{ChildLeaf, RegTree, SplitRule};
fn collapsed_tree() -> RegTree {
let mut t = RegTree::with_root(1.0);
let (l, _) = t.expand(
0,
SplitRule::numeric(0, 0.5, true),
ChildLeaf::new(0.0, 1.0),
ChildLeaf::new(7.0, 1.0),
);
let (ll, lr) = t.expand(
l,
SplitRule::numeric(1, 2.0, false),
ChildLeaf::new(0.0, 1.0),
ChildLeaf::new(0.0, 1.0),
);
t.expand(
ll,
SplitRule::numeric(2, -1.0, true),
ChildLeaf::new(-3.0, 1.0),
ChildLeaf::new(-2.0, 1.0),
);
t.expand(
lr,
SplitRule::numeric(2, -1.0, true),
ChildLeaf::new(1.5, 1.0),
ChildLeaf::new(2.5, 1.0),
);
t
}
fn asymmetric_tree() -> RegTree {
let mut t = RegTree::with_root(1.0);
let (l, r) = t.expand(
0,
SplitRule::numeric(0, 0.5, true),
ChildLeaf::new(0.0, 1.0),
ChildLeaf::new(0.0, 1.0),
);
t.expand(
l,
SplitRule::numeric(1, 2.0, false),
ChildLeaf::new(1.0, 1.0),
ChildLeaf::new(2.0, 1.0),
);
t.expand(
r,
SplitRule::numeric(1, 3.0, false),
ChildLeaf::new(3.0, 1.0),
ChildLeaf::new(4.0, 1.0),
);
t
}
fn chain_tree(depth: usize) -> RegTree {
let mut t = RegTree::with_root(1.0);
let mut node = 0;
for d in 0..depth {
let (l, _) = t.expand(
node,
SplitRule::numeric(0, d as f32, true),
ChildLeaf::new(0.0, 1.0),
ChildLeaf::new(d as f32, 1.0),
);
node = l;
}
t
}
fn rows(n: usize, n_cols: usize) -> Vec<f32> {
let specials = [
f32::NAN,
0.5,
0.5f32.next_down(),
0.5f32.next_up(),
2.0,
2.0f32.next_down(),
-1.0,
(-1.0f32).next_up(),
-0.0,
f32::INFINITY,
f32::NEG_INFINITY,
];
(0..n * n_cols)
.map(|i| {
let h = (i as u64).wrapping_mul(crate::rng::GOLDEN) >> 40;
if h.is_multiple_of(3) {
specials[(h / 3) as usize % specials.len()]
} else {
(h % 1000) as f32 / 200.0 - 2.0
}
})
.collect()
}
#[test]
fn detection_accepts_level_uniform_trees_only() {
let trees = [
collapsed_tree(),
asymmetric_tree(),
chain_tree(1),
chain_tree(4),
chain_tree(5),
];
let forest = CompactForest::from_trees(&trees);
let symmetric: Vec<bool> = (0..trees.len()).map(|t| forest.is_symmetric(t)).collect();
assert_eq!(symmetric, [true, false, false, true, false]);
}
#[test]
fn bit_pattern_walk_matches_reference_routing() {
let trees = [collapsed_tree(), chain_tree(4), asymmetric_tree()];
let forest = CompactForest::from_trees(&trees);
let n_cols = 3;
let n = 5 * LANES + 7;
let data = rows(n, n_cols);
let (lanes, tail) = split_lanes(&data, n_cols);
let block = LaneBlock {
lanes: &lanes,
tail,
n_cols,
rows: n,
};
for (t, tree) in trees.iter().enumerate() {
let mut margins = vec![0.25f32; n];
forest.accumulate(t, block, 1.0, &mut margins, 1);
let mut leaves = vec![0u32; n];
forest.original_leaf_ids(t, block, &mut leaves, 1);
for (r, row) in data.chunks_exact(n_cols).enumerate() {
let leaf = tree.leaf_id_dense(row, f32::NAN);
assert_eq!(leaves[r] as usize, leaf, "tree {t} row {r} {row:?}");
let want = 0.25f32 + tree.node(leaf).leaf_value;
assert_eq!(margins[r].to_bits(), want.to_bits());
}
}
}
#[test]
fn trained_symmetric_model_predicts_bit_identically() {
let (n, n_cols) = (3000, 5);
let x: Vec<f32> = rows(n, n_cols)
.into_iter()
.map(|v| if v.is_infinite() { v.signum() * 3.0 } else { v })
.collect();
let y: Vec<f32> = x
.chunks_exact(n_cols)
.map(|r| r.iter().filter(|v| !v.is_nan()).map(|v| v.sin()).sum())
.collect();
let data = labeled_dense(&x, n, n_cols, &y);
let params = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.grow_policy(GrowPolicy::Symmetric)
.max_depth(5)
.build()
.unwrap();
let model = train(¶ms, &data, 30).unwrap();
let forest = CompactForest::from_trees(model.trees());
assert!((0..model.num_trees()).all(|t| forest.is_symmetric(t)));
let margins = model.predict_margin(&data, Iterations::Best).unwrap();
let leaves = model.predict_leaf(&data, ..).unwrap();
for (r, row) in x.chunks_exact(n_cols).enumerate() {
let mut want = model.base_score();
for (t, tree) in model.trees().iter().enumerate() {
let leaf = tree.leaf_id_dense(row, f32::NAN);
assert_eq!(leaves.get(r, t).copied(), Some(leaf as u32));
want += 1.0 * tree.node(leaf).leaf_value;
}
assert_eq!(
margins.get(r, 0).map(|m| m.to_bits()),
Some(want.to_bits()),
"row {r}"
);
}
}
}