use crate::tree::RegTree;
pub(crate) const LANES: usize = 16;
pub(crate) const FEATURE_LANES: usize = 2 * LANES;
const MAX_FIXED_DEPTH: u32 = 16;
const MAX_SLOT_FEATURE: u32 = u32::MAX / FEATURE_LANES as u32;
const SIGN: u32 = 1 << 31;
#[inline(always)]
pub(crate) fn key(v: f32) -> u32 {
if v.is_nan() {
return 0;
}
let bits = (v + 0.0).to_bits();
bits ^ ((((bits as i32) >> 31) as u32) | SIGN)
}
#[inline(always)]
fn unkey(key: u32) -> f32 {
f32::from_bits(if key & SIGN != 0 { key ^ SIGN } else { !key })
}
const LEAF_KEY: u32 = 0xFF80_0000;
const _: () = assert!(LEAF_KEY == f32::INFINITY.to_bits() | SIGN);
const CATEGORICAL: u32 = 1;
const CAT_DEFAULT_LEFT: u32 = 2;
const CAT_END_SHIFT: u32 = 2;
#[derive(Debug, Clone, Copy)]
#[repr(C)]
struct CNode {
slot: u32,
key: u32,
left: u32,
aux: u32,
}
#[cfg(target_arch = "x86_64")]
const _: () = assert!(
std::mem::size_of::<CNode>() == 16
&& std::mem::offset_of!(CNode, slot) == 0
&& std::mem::offset_of!(CNode, key) == 4
);
impl CNode {
#[inline(always)]
fn feature(&self) -> usize {
self.slot as usize / FEATURE_LANES
}
#[inline(always)]
fn negate_mask(&self) -> u32 {
(self.slot & LANES as u32) << (31 - LANES.trailing_zeros())
}
}
fn next_below(c: f32) -> f32 {
debug_assert!(c.is_finite());
if c > 0.0 {
f32::from_bits(c.to_bits() - 1)
} else if c == 0.0 {
-f32::from_bits(1)
} else {
f32::from_bits(c.to_bits() + 1)
}
}
#[derive(Debug, Clone, Copy)]
struct TreeMeta {
root: u32,
depth: u32,
has_categorical: bool,
max_feature: u32,
}
impl TreeMeta {
#[inline]
fn lockstep_ok(&self) -> bool {
!self.has_categorical && self.depth <= MAX_FIXED_DEPTH
}
#[inline]
fn check_width(&self, n_cols: usize) {
assert!(
(self.max_feature as usize) < n_cols,
"row has {n_cols} features but the tree splits on feature {}",
self.max_feature
);
}
}
#[derive(Debug, Clone)]
pub(crate) struct CompactForest {
nodes: Vec<CNode>,
orig_id: Vec<u32>,
categories: Vec<u32>,
trees: Vec<TreeMeta>,
}
impl CompactForest {
pub(crate) fn from_trees(trees: &[RegTree]) -> Self {
let total: usize = trees.iter().map(RegTree::num_nodes).sum();
let mut forest = CompactForest {
nodes: Vec::with_capacity(total),
orig_id: Vec::with_capacity(total),
categories: Vec::new(),
trees: Vec::with_capacity(trees.len()),
};
for tree in trees {
forest.push_tree(tree);
}
forest
}
fn push_tree(&mut self, tree: &RegTree) {
let src = tree.nodes();
let base = self.nodes.len() as u32;
let cat_base = self.categories.len() as u32;
self.categories.extend_from_slice(tree.categories());
let mut order: Vec<u32> = Vec::with_capacity(src.len());
let mut new_id = vec![u32::MAX; src.len()];
let mut depth_of = vec![0u32; src.len()];
order.push(0);
new_id[0] = base;
let mut i = 0;
while i < order.len() {
let old = order[i] as usize;
let n = &src[old];
if !n.is_leaf() {
let (first, second) = if n.default_left || n.is_categorical {
(n.left, n.right)
} else {
(n.right, n.left)
};
for child in [first as usize, second as usize] {
new_id[child] = base + order.len() as u32;
depth_of[child] = depth_of[old] + 1;
order.push(child as u32);
}
}
i += 1;
}
let mut has_categorical = false;
let mut depth = 0u32;
let mut max_feature = 0u32;
for &old in &order {
let n = &src[old as usize];
let id = self.nodes.len() as u32;
depth = depth.max(depth_of[old as usize]);
let node = if n.is_leaf() {
CNode {
slot: 0,
key: LEAF_KEY,
left: id,
aux: n.leaf_value.to_bits(),
}
} else if n.is_categorical {
has_categorical = true;
max_feature = max_feature.max(n.split_feature);
let (begin, end) = (cat_base + n.cat_begin, cat_base + n.cat_end);
assert!(
end < (1 << (32 - CAT_END_SHIFT)),
"categorical set range does not fit the compact encoding"
);
let mut aux = CATEGORICAL | (end << CAT_END_SHIFT);
if n.default_left {
aux |= CAT_DEFAULT_LEFT;
}
CNode {
slot: n.split_feature * FEATURE_LANES as u32,
key: begin,
left: new_id[n.left as usize],
aux,
}
} else if n.default_left {
max_feature = max_feature.max(n.split_feature);
CNode {
slot: n.split_feature * FEATURE_LANES as u32,
key: key(next_below(n.split_cond)),
left: new_id[n.left as usize],
aux: 0,
}
} else {
max_feature = max_feature.max(n.split_feature);
CNode {
slot: n.split_feature * FEATURE_LANES as u32 + LANES as u32,
key: key(-n.split_cond),
left: new_id[n.right as usize],
aux: 0,
}
};
self.nodes.push(node);
self.orig_id.push(old);
}
assert!(
max_feature <= MAX_SLOT_FEATURE,
"feature index {max_feature} does not fit the compact split encoding"
);
self.trees.push(TreeMeta {
root: base,
depth,
has_categorical,
max_feature,
});
}
#[inline]
pub(crate) fn leaf_value(&self, id: u32) -> f32 {
f32::from_bits(self.nodes[id as usize].aux)
}
#[inline]
pub(crate) fn original_id(&self, id: u32) -> u32 {
self.orig_id[id as usize]
}
#[inline]
fn is_leaf(node: &CNode, nid: u32) -> bool {
node.left == nid
}
#[inline]
fn in_left_set(&self, node: &CNode, v: f32) -> bool {
let begin = node.key as usize;
let end = (node.aux >> CAT_END_SHIFT) as usize;
let c = v as u32;
self.categories[begin..end].contains(&c)
}
#[inline]
fn next(&self, node: &CNode, v: f32) -> u32 {
if node.aux & CATEGORICAL != 0 {
let go_left = if v.is_nan() {
node.aux & CAT_DEFAULT_LEFT != 0
} else {
self.in_left_set(node, v)
};
node.left + u32::from(!go_left)
} else {
let v = f32::from_bits(v.to_bits() ^ node.negate_mask());
Self::next_numeric(node, key(v)) as u32
}
}
#[inline(always)]
fn next_numeric(node: &CNode, key: u32) -> usize {
node.left as usize + usize::from(key > node.key)
}
#[inline(always)]
fn slot_key(node: &CNode) -> (usize, u32) {
#[cfg(target_arch = "x86_64")]
{
let packed = unsafe {
std::ptr::from_ref::<CNode>(node)
.cast::<u64>()
.read_unaligned()
};
(packed as u32 as usize, (packed >> 32) as u32)
}
#[cfg(not(target_arch = "x86_64"))]
{
(node.slot as usize, node.key)
}
}
#[inline]
pub(crate) fn leaf_id(&self, t: usize, row: &[f32]) -> u32 {
let mut nid = self.trees[t].root;
loop {
let node = &self.nodes[nid as usize];
if Self::is_leaf(node, nid) {
return nid;
}
nid = self.next(node, row[node.feature()]);
}
}
pub(crate) fn leaf_id_with(&self, t: usize, get: impl Fn(u32) -> Option<f32>) -> u32 {
let mut nid = self.trees[t].root;
loop {
let node = &self.nodes[nid as usize];
if Self::is_leaf(node, nid) {
return nid;
}
let v = get(node.feature() as u32).unwrap_or(f32::NAN);
nid = self.next(node, v);
}
}
#[inline]
fn leaf_id_keyed(&self, t: usize, grp: &[u32], lane: usize) -> u32 {
let mut nid = self.trees[t].root;
loop {
let node = &self.nodes[nid as usize];
if Self::is_leaf(node, nid) {
return nid;
}
nid = if node.aux & CATEGORICAL != 0 {
self.next(node, unkey(grp[node.slot as usize + lane]))
} else {
Self::next_numeric(node, grp[node.slot as usize + lane]) as u32
};
}
}
#[inline(always)]
fn walk_block(
&self,
t: usize,
lanes: &[u32],
tail: &[f32],
n_cols: usize,
rows: usize,
mut sink: impl FnMut(usize, u32),
) {
let groups = rows / LANES;
let group_len = FEATURE_LANES * n_cols;
assert!(
lanes.len() >= groups * group_len && tail.len() >= (rows - groups * LANES) * n_cols,
"row block holds fewer rows than requested"
);
let meta = self.trees[t];
if meta.lockstep_ok() {
meta.check_width(n_cols);
let nodes = &self.nodes[..];
let root = meta.root as usize;
for g in 0..groups {
let grp = &lanes[g * group_len..(g + 1) * group_len];
let mut nid = [root; LANES];
for _ in 0..meta.depth {
macro_rules! lane {
($($j:literal)*) => {$(
let node = unsafe { nodes.get_unchecked(nid[$j]) };
let (slot, key) = Self::slot_key(node);
let k = unsafe { *grp.get_unchecked(slot + $j) };
nid[$j] = node.left as usize + usize::from(k > key);
)*};
}
lane!(0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15);
}
for (j, &id) in nid.iter().enumerate() {
sink(g * LANES + j, id as u32);
}
}
} else {
for g in 0..groups {
let grp = &lanes[g * group_len..(g + 1) * group_len];
for j in 0..LANES {
sink(g * LANES + j, self.leaf_id_keyed(t, grp, j));
}
}
}
for (i, row) in tail
.chunks_exact(n_cols)
.take(rows - groups * LANES)
.enumerate()
{
sink(groups * LANES + i, self.leaf_id(t, row));
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn original_leaf_ids(
&self,
t: usize,
lanes: &[u32],
tail: &[f32],
n_cols: usize,
rows: usize,
out: &mut [u32],
stride: usize,
) {
assert!(rows == 0 || out.len() > (rows - 1) * stride);
let orig = &self.orig_id[..];
self.walk_block(t, lanes, tail, n_cols, rows, |r, leaf| {
unsafe {
*out.get_unchecked_mut(r * stride) = *orig.get_unchecked(leaf as usize);
}
});
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn accumulate(
&self,
t: usize,
lanes: &[u32],
tail: &[f32],
n_cols: usize,
rows: usize,
weight: f32,
out: &mut [f32],
stride: usize,
) {
assert!(rows == 0 || out.len() > (rows - 1) * stride);
let nodes = &self.nodes[..];
self.walk_block(t, lanes, tail, n_cols, rows, |r, leaf| {
unsafe {
*out.get_unchecked_mut(r * stride) +=
weight * f32::from_bits(nodes.get_unchecked(leaf as usize).aux);
}
});
}
#[inline(always)]
fn walk_row(&self, row: &[f32], limit: usize, mut sink: impl FnMut(usize, u32)) {
assert!(limit <= self.trees.len());
let nodes = &self.nodes[..];
let full = limit / LANES * LANES;
let mut keys: Vec<u32> = Vec::new();
for g in 0..limit / LANES {
let group = &self.trees[g * LANES..(g + 1) * LANES];
let mut depth = 0u32;
let mut ok = true;
for meta in group {
ok &= meta.lockstep_ok();
depth = depth.max(meta.depth);
}
if !ok {
for j in 0..LANES {
sink(g * LANES + j, self.leaf_id(g * LANES + j, row));
}
continue;
}
for meta in group {
meta.check_width(row.len());
}
if keys.is_empty() {
keys.reserve(2 * row.len());
for &v in row {
keys.push(key(v));
keys.push(key(-v));
}
}
let keys = &keys[..];
let mut nid = [0usize; LANES];
for (n, meta) in nid.iter_mut().zip(group) {
*n = meta.root as usize;
}
for _ in 0..depth {
macro_rules! lane {
($($j:literal)*) => {$(
let node = unsafe { nodes.get_unchecked(nid[$j]) };
let k = unsafe { *keys.get_unchecked(node.slot as usize / LANES) };
nid[$j] = Self::next_numeric(node, k);
)*};
}
lane!(0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15);
}
for (j, &id) in nid.iter().enumerate() {
sink(g * LANES + j, id as u32);
}
}
for t in full..limit {
sink(t, self.leaf_id(t, row));
}
}
pub(crate) fn original_leaf_ids_for_row(&self, row: &[f32], out: &mut [u32]) {
let orig = &self.orig_id[..];
self.walk_row(row, out.len(), |t, leaf| out[t] = orig[leaf as usize]);
}
pub(crate) fn accumulate_row(
&self,
row: &[f32],
limit: usize,
weight: impl Fn(usize) -> f32,
out: &mut [f32],
) {
let k = out.len();
let nodes = &self.nodes[..];
self.walk_row(row, limit, |t, leaf| {
out[t % k] += weight(t) * f32::from_bits(nodes[leaf as usize].aux);
});
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tree() -> RegTree {
let mut t = RegTree::with_root(1.0);
let (l, _r) = t.expand(0, 0, 0.5, true, 0.0, 1.0, 5.0, 1.0);
t.expand(l, 1, 2.0, false, -1.0, 1.0, 1.0, 1.0);
t
}
fn split_lanes(rows: &[f32], n_cols: usize) -> (Vec<u32>, &[f32]) {
let groups = rows.len() / n_cols / LANES;
let mut lanes = vec![0u32; groups * FEATURE_LANES * n_cols];
for g in 0..groups {
for j in 0..LANES {
for f in 0..n_cols {
let v = rows[(g * LANES + j) * n_cols + f];
let base = g * FEATURE_LANES * n_cols + f * FEATURE_LANES + j;
lanes[base] = key(v);
lanes[base + LANES] = key(-v);
}
}
}
(lanes, &rows[groups * LANES * n_cols..])
}
#[test]
fn matches_reference_walk() {
let t = tree();
let f = CompactForest::from_trees(std::slice::from_ref(&t));
assert_eq!(f.trees[0].depth, 2);
let rows: Vec<[f32; 2]> = vec![
[0.1, 1.0],
[0.1, 3.0],
[0.9, 0.0],
[f32::NAN, 1.0],
[f32::NAN, f32::NAN],
[0.1, f32::NAN],
];
let flat: Vec<f32> = rows.iter().flatten().copied().collect();
let mut out = vec![0u32; rows.len()];
f.original_leaf_ids(0, &[], &flat, 2, rows.len(), &mut out, 1);
for (r, row) in rows.iter().enumerate() {
let want = t.leaf_id_dense(row, f32::NAN);
assert_eq!(out[r] as usize, want, "row {r}");
let leaf = f.leaf_id(0, row);
assert_eq!(f.original_id(leaf) as usize, want);
assert_eq!(f.leaf_value(leaf), t.node(want).leaf_value);
}
}
#[test]
fn lockstep_matches_scalar_on_full_blocks() {
let t = tree();
let f = CompactForest::from_trees(&[t]);
let n = 3 * LANES + 5;
let block: Vec<f32> = (0..n * 2)
.map(|i| {
if i % 7 == 0 {
f32::NAN
} else {
(i % 5) as f32 * 0.3
}
})
.collect();
let (lanes, tail) = split_lanes(&block, 2);
let mut ids = vec![0u32; n * 3];
f.original_leaf_ids(0, &lanes, tail, 2, n, &mut ids[1..], 3);
let mut acc = vec![0.5f32; n * 2];
f.accumulate(0, &lanes, tail, 2, n, 2.0, &mut acc, 2);
for r in 0..n {
let leaf = f.leaf_id(0, &block[r * 2..r * 2 + 2]);
assert_eq!(ids[r * 3 + 1], f.original_id(leaf));
assert_eq!(acc[r * 2], 0.5 + 2.0 * f.leaf_value(leaf));
assert_eq!(acc[r * 2 + 1], 0.5);
}
}
#[test]
fn tree_lockstep_matches_per_tree_walk() {
let mut trees: Vec<RegTree> = (0..2 * LANES + 3)
.map(|i| {
let mut t = RegTree::with_root(1.0);
let (l, _r) = t.expand(
0,
i as u32 % 3,
0.1 * i as f32,
i % 2 == 0,
0.0,
1.0,
1.0,
1.0,
);
t.expand(l, 1, 0.5, i % 4 == 0, -1.0, 1.0, 2.0, 1.0);
t
})
.collect();
let mut cat = RegTree::with_root(1.0);
cat.expand_categorical(0, 2, &[1, 3], false, -1.0, 1.0, 1.0, 1.0);
trees[LANES + 1] = cat;
let f = CompactForest::from_trees(&trees);
let row = [0.7f32, f32::NAN, 3.0];
let mut out = vec![0u32; trees.len()];
f.original_leaf_ids_for_row(&row, &mut out);
let mut acc = vec![0.25f32; 3];
f.accumulate_row(&row, trees.len(), |t| 1.0 + t as f32, &mut acc);
let mut want_acc = vec![0.25f32; 3];
for (t, tree) in trees.iter().enumerate() {
let want = tree.leaf_id_dense(&row, f32::NAN);
assert_eq!(out[t] as usize, want, "tree {t}");
want_acc[t % 3] += (1.0 + t as f32) * tree.node(want).leaf_value;
}
assert_eq!(acc, want_acc);
}
#[test]
fn categorical_membership() {
for default_left in [true, false] {
let mut t = RegTree::with_root(1.0);
t.expand_categorical(0, 0, &[2, 5], default_left, -1.0, 1.0, 1.0, 1.0);
let mut first = RegTree::with_root(1.0);
first.expand_categorical(0, 0, &[9], true, 0.0, 1.0, 0.0, 1.0);
let f = CompactForest::from_trees(&[first, t.clone()]);
assert!(f.trees[1].has_categorical);
for v in [0.0f32, 2.0, 5.0, 7.0, f32::NAN] {
let want = t.leaf_id_dense(&[v], f32::NAN);
let got = f.leaf_id(1, &[v]);
assert_eq!(f.original_id(got) as usize, want, "v={v} dl={default_left}");
assert_eq!(f.leaf_value(got), t.node(want).leaf_value);
}
}
}
#[test]
fn keys_order_like_floats_and_isolate_missing() {
let values = [
f32::NEG_INFINITY,
-f32::MAX,
-1.5,
-f32::MIN_POSITIVE,
-0.0,
0.0,
f32::from_bits(1),
2.0,
f32::MAX,
f32::INFINITY,
];
for (i, &a) in values.iter().enumerate() {
for &b in &values[i..] {
assert_eq!(key(a) > key(b), a > b, "{a} vs {b}");
assert_eq!(key(a) == key(b), a == b, "{a} vs {b}");
}
assert!(key(a) > key(f32::NAN), "{a} must key above missing");
assert!((key(f32::NAN) <= key(a)), "missing never compares greater");
assert!(unkey(key(a)) == a, "{a} round trip");
}
assert_eq!(key(f32::NAN), 0);
assert_eq!(key(-f32::NAN), 0);
assert!(unkey(0).is_nan());
assert_eq!(key(f32::INFINITY), LEAF_KEY);
}
#[test]
fn boundary_values_match_reference_in_every_path() {
let mut trees = Vec::new();
for (cond, default_left) in [
(0.0f32, true),
(0.0, false),
(-1.5, true),
(-1.5, false),
(f32::MAX, true),
(-f32::MAX, false),
(f32::MIN_POSITIVE, false),
] {
let mut t = RegTree::with_root(1.0);
let (l, r) = t.expand(0, 0, cond, default_left, 0.0, 1.0, 0.0, 1.0);
t.expand(l, 1, cond, !default_left, -1.0, 1.0, 1.0, 1.0);
t.expand(r, 1, -cond, default_left, 2.0, 1.0, 3.0, 1.0);
trees.push(t);
}
let f = CompactForest::from_trees(&trees);
let probes = [
f32::NEG_INFINITY,
-f32::MAX,
-1.5,
-f32::from_bits(f32::MIN_POSITIVE.to_bits() + 1),
-f32::MIN_POSITIVE,
-0.0,
0.0,
f32::MIN_POSITIVE,
1.5,
f32::MAX,
f32::INFINITY,
f32::NAN,
];
let mut rows: Vec<f32> = vec![0.5, -0.5];
for &a in &probes {
for &b in &probes {
rows.extend_from_slice(&[a, b]);
}
}
let n = rows.len() / 2;
assert!(!n.is_multiple_of(LANES), "layout must exercise a tail");
let (lanes, tail) = split_lanes(&rows, 2);
for (t, tree) in trees.iter().enumerate() {
let mut ids = vec![0u32; n];
f.original_leaf_ids(t, &lanes, tail, 2, n, &mut ids, 1);
for r in 0..n {
let row = &rows[r * 2..r * 2 + 2];
let want = tree.leaf_id_dense(row, f32::NAN);
assert_eq!(ids[r] as usize, want, "tree {t} row {row:?} (block)");
let leaf = f.leaf_id(t, row);
assert_eq!(f.original_id(leaf) as usize, want, "tree {t} row {row:?}");
let mut per_tree = vec![0u32; trees.len()];
f.original_leaf_ids_for_row(row, &mut per_tree);
assert_eq!(per_tree[t] as usize, want, "tree {t} row {row:?} (row)");
}
}
}
#[test]
fn infinite_mirrored_threshold_is_not_mistaken_for_a_leaf() {
let mut t = RegTree::with_root(1.0);
t.expand(0, 0, f32::NEG_INFINITY, false, -1.0, 1.0, 1.0, 1.0);
let f = CompactForest::from_trees(std::slice::from_ref(&t));
for v in [f32::NEG_INFINITY, -1.0, 0.0, 1.0, f32::INFINITY, f32::NAN] {
let want = t.leaf_id_dense(&[v], f32::NAN);
assert_eq!(f.original_id(f.leaf_id(0, &[v])) as usize, want, "v={v}");
let mut out = [0u32];
f.original_leaf_ids_for_row(&[v], &mut out);
assert_eq!(out[0] as usize, want, "v={v} (row)");
}
}
}