use crate::data::DMatrix;
use crate::data::meta::FeatureType;
use crate::data::sketch::{SketchScratch, WQSketch};
use crate::data::sort::{RadixScratch, sort_values};
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>,
}
#[inline]
fn global_bin(start: usize, local: usize, n_cuts: usize) -> u32 {
start as u32 + local.min(n_cuts.saturating_sub(1)) as u32
}
#[derive(Clone, Copy)]
struct Ingest {
sorted: bool,
unit_weights: bool,
}
impl HistCuts {
pub fn from_dmatrix(data: &DMatrix, max_bin: usize) -> Self {
let weights = data.weights();
let unit_weights = weights.is_none();
Self::build(
data,
max_bin,
|row| weights.map_or(1.0, |w| w[row]),
Ingest {
sorted: false,
unit_weights,
},
)
}
pub fn from_dmatrix_weighted(
data: &DMatrix,
max_bin: usize,
hessians: &[f32],
sorted: bool,
) -> Self {
Self::from_dmatrix_hessians(data, max_bin, |row| hessians[row], sorted)
}
pub(crate) fn from_dmatrix_hessians(
data: &DMatrix,
max_bin: usize,
hessian: impl Fn(usize) -> f32 + Sync,
sorted: bool,
) -> Self {
let weights = data.weights();
Self::build(
data,
max_bin,
|row| weights.map_or(hessian(row), |w| hessian(row) * w[row]),
Ingest {
sorted,
unit_weights: false,
},
)
}
fn build(
data: &DMatrix,
max_bin: usize,
weight_of: impl Fn(usize) -> f32 + Sync,
ingest: Ingest,
) -> Self {
let Ingest {
sorted,
unit_weights,
} = ingest;
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 is_categorical: Vec<bool> = ftypes
.iter()
.map(|&t| t == FeatureType::Categorical)
.collect();
let mut feature_offset = Vec::with_capacity(n_features + 1);
feature_offset.push(0u32);
let mut cut_values = Vec::new();
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 ColumnScratch, output: &mut Vec<f32>| {
let ColumnScratch {
values,
sort: spare,
pairs,
sketch_sort,
} = 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 n_values = csc.as_ref().map_or_else(
|| {
let mut n = 0usize;
for_each_value(f, &mut |_, _| n += 1);
n
},
|csc| csc.col_len(f),
);
let mut sketch =
WQSketch::new(n_values, max_bin).with_sort_scratch(std::mem::take(sketch_sort));
if unit_weights {
sketch = sketch.with_unit_weights();
}
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);
*sketch_sort = sketch.into_sort_scratch();
};
if n_features > 1
&& n_rows.saturating_mul(n_features) >= 65_536
&& rayon::current_num_threads() > 1
{
let columns: Vec<_> = (0..n_features)
.into_par_iter()
.map_init(ColumnScratch::default, |scratch, f| {
let mut output = Vec::new();
build(f, scratch, &mut output);
output
})
.collect();
for column in columns {
cut_values.extend(column);
feature_offset.push(cut_values.len() as u32);
}
} else {
let mut scratch = ColumnScratch::default();
for f in 0..n_features {
build(f, &mut scratch, &mut cut_values);
feature_offset.push(cut_values.len() as u32);
}
}
HistCuts {
n_features,
feature_offset,
cut_values,
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,
features: Vec<FeatureSearch>,
padded: Vec<f32>,
level1: Vec<f32>,
}
#[derive(Debug, Clone, Copy)]
struct FeatureSearch {
start: usize,
n_cuts: usize,
padded: usize,
padded_len: usize,
level1: usize,
level1_len: usize,
categorical: bool,
}
impl<'a> BinSearch<'a> {
pub fn new(cuts: &'a HistCuts) -> Self {
let mut padded = Vec::new();
let mut level1 = Vec::new();
let mut features = Vec::with_capacity(cuts.n_features());
for f in 0..cuts.n_features() {
let (start, end) = cuts.feature_bins(f);
let (padded_start, level1_start) = (padded.len(), level1.len());
let categorical = cuts.is_categorical(f);
if !categorical {
let feature = &cuts.cut_values[start..end];
let blocks = feature.len().div_ceil(SEARCH_BLOCK);
padded.extend_from_slice(feature);
padded.resize(padded_start + blocks * SEARCH_BLOCK, f32::INFINITY);
level1.extend(
feature
.chunks(SEARCH_BLOCK)
.map(|block| *block.last().expect("blocks are non-empty")),
);
level1.resize(
level1_start + blocks.div_ceil(SEARCH_BLOCK) * SEARCH_BLOCK,
f32::INFINITY,
);
}
features.push(FeatureSearch {
start,
n_cuts: end - start,
padded: padded_start,
padded_len: padded.len() - padded_start,
level1: level1_start,
level1_len: level1.len() - level1_start,
categorical,
});
}
BinSearch {
cuts,
features,
padded,
level1,
}
}
#[inline]
pub fn n_features(&self) -> usize {
self.cuts.n_features()
}
#[inline(always)]
pub fn bin_of(&self, f: usize, value: f32) -> u32 {
self.feature(f).bin_of(value)
}
#[inline(always)]
pub(crate) fn feature(&self, f: usize) -> FeatureBins<'_> {
let feature = self.features[f];
FeatureBins {
cuts: self.cuts,
f,
categorical: feature.categorical,
start: feature.start,
n_cuts: feature.n_cuts,
level1: &self.level1[feature.level1..][..feature.level1_len],
padded: &self.padded[feature.padded..][..feature.padded_len],
}
}
}
#[derive(Clone, Copy)]
pub(crate) struct FeatureBins<'a> {
cuts: &'a HistCuts,
f: usize,
categorical: bool,
start: usize,
n_cuts: usize,
level1: &'a [f32],
padded: &'a [f32],
}
impl FeatureBins<'_> {
#[inline(always)]
pub(crate) fn bin_of(&self, value: f32) -> u32 {
if self.categorical {
return self.cuts.bin_of(self.f, value);
}
let mut block = 0;
for chunk in self.level1.as_chunks::<SEARCH_BLOCK>().0 {
block += crate::simd::count_le(chunk, value);
}
let local = if block * SEARCH_BLOCK < self.padded.len() {
block * SEARCH_BLOCK
+ crate::simd::count_le(
&self.padded[block * SEARCH_BLOCK..(block + 1) * SEARCH_BLOCK],
value,
)
} else {
self.n_cuts
};
global_bin(self.start, local, self.n_cuts)
}
}
#[derive(Default)]
struct ColumnScratch {
values: Vec<f32>,
sort: RadixScratch<f32>,
pairs: Vec<(f32, f32)>,
sketch_sort: SketchScratch,
}
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 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_has_one_bin() {
let data = DMatrix::from_dense(&[5.0, 5.0, 5.0], 3, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
assert_eq!(cuts.num_bins(0), 1);
assert_eq!(cuts.bin_of(0, 5.0), 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}"
);
}
}
}
}