use crate::data::DMatrix;
use crate::data::meta::FeatureType;
use crate::data::sketch::WQSketch;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HistCuts {
n_features: usize,
feature_offset: Vec<u32>,
cut_values: Vec<f32>,
is_categorical: Vec<bool>,
}
fn is_cat(ftypes: &[FeatureType], f: usize) -> bool {
ftypes.get(f).copied() == Some(FeatureType::Categorical)
}
struct CutAssembler {
feature_offset: Vec<u32>,
cut_values: Vec<f32>,
}
impl CutAssembler {
fn new(n_features: usize) -> Self {
let mut feature_offset = Vec::with_capacity(n_features + 1);
feature_offset.push(0u32);
CutAssembler {
feature_offset,
cut_values: Vec::new(),
}
}
fn finish_feature(&mut self) {
self.feature_offset.push(self.cut_values.len() as u32);
}
fn assemble(self, n_features: usize, is_categorical: Vec<bool>) -> HistCuts {
HistCuts {
n_features,
feature_offset: self.feature_offset,
cut_values: self.cut_values,
is_categorical,
}
}
}
#[inline]
fn global_bin(start: usize, local: usize, n_cuts: usize) -> u32 {
start as u32 + local.min(n_cuts.saturating_sub(1)) as u32
}
impl HistCuts {
pub fn from_dmatrix(data: &DMatrix, max_bin: usize) -> Self {
let weights = data.weights();
Self::build(data, max_bin, |row| weights.map_or(1.0, |w| w[row]), false)
}
pub fn from_dmatrix_weighted(
data: &DMatrix,
max_bin: usize,
hessians: &[f32],
sorted: bool,
) -> Self {
let weights = data.weights();
Self::build(
data,
max_bin,
|row| weights.map_or(hessians[row], |w| hessians[row] * w[row]),
sorted,
)
}
fn build(
data: &DMatrix,
max_bin: usize,
weight_of: impl Fn(usize) -> f32 + Sync,
sorted: bool,
) -> Self {
let n_rows = data.n_rows();
let n_features = data.n_cols();
let (columns, csc) = match data.dense_values() {
Some(values) => (
Some(crate::data::ghist::transpose_dense(
values, n_rows, n_features,
)),
None,
),
None => (None, Some(data.to_csc())),
};
let missing = data.missing();
let ftypes = data.feature_types();
let mut assembler = CutAssembler::new(n_features);
let is_categorical: Vec<bool> = (0..n_features).map(|f| is_cat(ftypes, f)).collect();
let for_each_value = |f: usize, visit: &mut dyn FnMut(usize, f32)| match (&columns, &csc) {
(Some(columns), _) => {
for (row, &v) in columns[f * n_rows..(f + 1) * n_rows].iter().enumerate() {
if !crate::data::dmatrix::is_missing(v, missing) {
visit(row, v);
}
}
}
(None, Some(csc)) => {
let (rows, values) = csc.column(f);
for (&row, &v) in rows.iter().zip(values) {
visit(row as usize, v);
}
}
(None, None) => unreachable!("one column source is always built"),
};
let build =
|f, scratch: &mut (Vec<f32>, Vec<f32>, Vec<(f32, f32)>), output: &mut Vec<f32>| {
let (values, spare, pairs) = scratch;
if is_categorical[f] {
values.clear();
for_each_value(f, &mut |_, v| values.push(v));
sort_values(values, spare);
build_categorical_cuts(values, output);
return;
}
let mut n_values = 0usize;
for_each_value(f, &mut |_, _| n_values += 1);
let mut sketch = WQSketch::new(n_values, max_bin);
if sorted {
pairs.clear();
for_each_value(f, &mut |row, v| pairs.push((v, weight_of(row))));
pairs.sort_by(|a, b| a.0.total_cmp(&b.0));
sketch.push_sorted(pairs);
} else {
for_each_value(f, &mut |row, v| sketch.push(v, weight_of(row)));
}
sketch.cut_values(output);
};
if n_features > 1
&& data.n_rows().saturating_mul(n_features) >= 65_536
&& rayon::current_num_threads() > 1
{
let columns: Vec<_> = (0..n_features)
.into_par_iter()
.map_init(Default::default, |scratch, f| {
let mut output = Vec::new();
build(f, scratch, &mut output);
output
})
.collect();
for column in columns {
assembler.cut_values.extend(column);
assembler.finish_feature();
}
} else {
let mut scratch = Default::default();
for f in 0..n_features {
build(f, &mut scratch, &mut assembler.cut_values);
assembler.finish_feature();
}
}
assembler.assemble(n_features, is_categorical)
}
#[inline]
pub fn is_categorical(&self, f: usize) -> bool {
self.is_categorical[f]
}
#[inline]
pub fn n_features(&self) -> usize {
self.n_features
}
#[inline]
pub fn total_bins(&self) -> usize {
self.cut_values.len()
}
#[inline]
pub fn feature_bins(&self, f: usize) -> (usize, usize) {
(
self.feature_offset[f] as usize,
self.feature_offset[f + 1] as usize,
)
}
#[inline]
pub fn num_bins(&self, f: usize) -> usize {
(self.feature_offset[f + 1] - self.feature_offset[f]) as usize
}
#[inline]
pub fn cut_value(&self, global_bin: usize) -> f32 {
self.cut_values[global_bin]
}
#[inline]
pub fn bin_of(&self, f: usize, value: f32) -> u32 {
let (start, end) = self.feature_bins(f);
let slice = &self.cut_values[start..end];
if self.is_categorical[f] {
let local = slice
.binary_search_by(|c| c.partial_cmp(&value).unwrap())
.unwrap_or(0);
return start as u32 + local as u32;
}
global_bin(start, slice.partition_point(|&c| c <= value), slice.len())
}
}
const SEARCH_BLOCK: usize = 16;
pub struct BinSearch<'a> {
cuts: &'a HistCuts,
padded: Vec<f32>,
padded_offset: Vec<usize>,
level1: Vec<f32>,
level1_offset: Vec<usize>,
}
impl<'a> BinSearch<'a> {
pub fn new(cuts: &'a HistCuts) -> Self {
let mut padded = Vec::new();
let mut padded_offset = vec![0];
let mut level1 = Vec::new();
let mut level1_offset = vec![0];
for f in 0..cuts.n_features() {
if !cuts.is_categorical(f) {
let (start, end) = cuts.feature_bins(f);
let feature = &cuts.cut_values[start..end];
let blocks = feature.len().div_ceil(SEARCH_BLOCK);
padded.extend_from_slice(feature);
padded.resize(padded_offset[f] + blocks * SEARCH_BLOCK, f32::INFINITY);
level1.extend(
feature
.chunks(SEARCH_BLOCK)
.map(|block| *block.last().expect("blocks are non-empty")),
);
level1.resize(
level1_offset[f] + blocks.div_ceil(SEARCH_BLOCK) * SEARCH_BLOCK,
f32::INFINITY,
);
}
padded_offset.push(padded.len());
level1_offset.push(level1.len());
}
BinSearch {
cuts,
padded,
padded_offset,
level1,
level1_offset,
}
}
#[inline]
pub fn n_features(&self) -> usize {
self.cuts.n_features()
}
#[inline]
pub fn bin_of(&self, f: usize, value: f32) -> u32 {
if self.cuts.is_categorical(f) {
return self.cuts.bin_of(f, value);
}
let (start, end) = self.cuts.feature_bins(f);
let level1 = &self.level1[self.level1_offset[f]..self.level1_offset[f + 1]];
let mut block = 0;
for chunk in level1.as_chunks::<SEARCH_BLOCK>().0 {
block += crate::simd::count_le(chunk, value);
}
let padded = &self.padded[self.padded_offset[f]..self.padded_offset[f + 1]];
let local = if block * SEARCH_BLOCK < padded.len() {
block * SEARCH_BLOCK
+ crate::simd::count_le(
&padded[block * SEARCH_BLOCK..(block + 1) * SEARCH_BLOCK],
value,
)
} else {
end - start
};
global_bin(start, local, end - start)
}
}
const RADIX_MIN_LEN: usize = 2048;
const RADIX_BITS: u32 = 11;
const RADIX_BUCKETS: usize = 1 << RADIX_BITS;
#[inline]
fn sort_key(value: f32) -> u32 {
let bits = value.to_bits();
if bits & 0x8000_0000 != 0 {
!bits
} else {
bits | 0x8000_0000
}
}
fn sort_values(values: &mut Vec<f32>, spare: &mut Vec<f32>) {
let n = values.len();
if n < RADIX_MIN_LEN {
values.sort_unstable_by(f32::total_cmp);
return;
}
let mut counts = vec![[0u32; RADIX_BUCKETS]; 3];
for &v in values.iter() {
let key = sort_key(v);
for (pass, count) in counts.iter_mut().enumerate() {
count[((key >> (RADIX_BITS * pass as u32)) & (RADIX_BUCKETS as u32 - 1)) as usize] += 1;
}
}
spare.clear();
spare.resize(n, 0.0);
let mut in_spare = false;
for (pass, count) in counts.iter_mut().enumerate() {
if count.iter().any(|&c| c as usize == n) {
continue;
}
let mut offset = 0u32;
for c in count.iter_mut() {
let start = offset;
offset += *c;
*c = start;
}
let shift = RADIX_BITS * pass as u32;
let (src, dst) = if in_spare {
(&*spare, &mut *values)
} else {
(&*values, &mut *spare)
};
for &v in src {
let bucket = ((sort_key(v) >> shift) & (RADIX_BUCKETS as u32 - 1)) as usize;
dst[count[bucket] as usize] = v;
count[bucket] += 1;
}
in_spare = !in_spare;
}
if in_spare {
std::mem::swap(values, spare);
}
}
fn build_categorical_cuts(sorted_vals: &[f32], out: &mut Vec<f32>) {
if sorted_vals.is_empty() {
out.push(0.0);
return;
}
out.push(sorted_vals[0]);
for w in sorted_vals.windows(2) {
if w[0] != w[1] {
out.push(w[1]);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn two_level_search_matches_bin_of() {
let n = 5000;
let mut x = vec![0f32; n * 3];
for r in 0..n {
x[r * 3] = ((r * 7919) % n) as f32 / n as f32;
x[r * 3 + 1] = (r % 3) as f32;
x[r * 3 + 2] = 2.5;
}
let data = DMatrix::from_dense(&x, n, 3).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
let search = BinSearch::new(&cuts);
for f in 0..3 {
let (start, end) = cuts.feature_bins(f);
let mut probes: Vec<f32> = cuts.cut_values[start..end].to_vec();
probes.extend(
cuts.cut_values[start..end]
.windows(2)
.map(|w| f32::midpoint(w[0], w[1])),
);
probes.extend([-1e9, 1e9, -0.0, 0.0, 0.5, 2.5, 3.0]);
for value in probes {
assert_eq!(
search.bin_of(f, value),
cuts.bin_of(f, value),
"feature {f} value {value}"
);
}
}
}
#[test]
fn radix_sort_matches_total_order() {
let mut seed = 0x9E37_79B9_7F4A_7C15u64;
let mut next = || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
seed
};
for n in [RADIX_MIN_LEN, RADIX_MIN_LEN + 1, 10_007, 65_536] {
let mut values: Vec<f32> = (0..n)
.map(|i| match i % 11 {
0 => -0.0,
1 => 0.0,
2 => f32::MAX,
3 => f32::MIN,
4 => f32::MIN_POSITIVE,
5 => -f32::MIN_POSITIVE,
_ => (next() as f32 / u64::MAX as f32 - 0.5) * 1e6,
})
.collect();
let mut expected = values.clone();
expected.sort_unstable_by(f32::total_cmp);
let mut spare = Vec::new();
sort_values(&mut values, &mut spare);
let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
assert_eq!(bits(&values), bits(&expected));
}
}
#[test]
fn few_distinct_values_one_bin_each() {
let data = DMatrix::from_dense(&[0.0, 1.0, 2.0, 1.0], 4, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
assert_eq!(cuts.n_features(), 1);
assert_eq!(cuts.num_bins(0), 3);
assert_eq!(
(
cuts.bin_of(0, 0.0),
cuts.bin_of(0, 1.0),
cuts.bin_of(0, 2.0)
),
(0, 1, 2)
);
}
#[test]
fn monotone_binning() {
let n = 1000;
let x: Vec<f32> = (0..n).map(|i| i as f32).collect();
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 16);
assert!(
cuts.num_bins(0) <= 16,
"at most max_bin - 1 cuts plus sentinel"
);
let mut prev = 0u32;
for i in 0..n {
let b = cuts.bin_of(0, i as f32);
assert!(b >= prev);
prev = b;
}
assert!(cuts.bin_of(0, 0.0) < cuts.bin_of(0, 999.0));
}
#[test]
fn constant_feature_never_collides() {
let data = DMatrix::from_dense(&[5.0, 5.0, 5.0], 3, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
assert_eq!(cuts.bin_of(0, 5.0), cuts.bin_of(0, 5.0));
}
#[test]
fn split_threshold_consistency() {
let x: Vec<f32> = (0..10).map(|i| i as f32).collect();
let data = DMatrix::from_dense(&x, 10, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
let (start, end) = cuts.feature_bins(0);
for target in start..end - 1 {
let thr = cuts.cut_value(target);
for &v in &x {
let goes_left_by_value = v < thr;
let goes_left_by_bin = cuts.bin_of(0, v) as usize <= target;
assert_eq!(
goes_left_by_value, goes_left_by_bin,
"value {v}, target bin {target}, thr {thr}"
);
}
}
}
}