use crate::data::quantile::{BinSearch, HistCuts};
use crate::data::{DMatrix, Entry};
use rayon::prelude::*;
use std::ops::Range;
#[derive(Debug, Clone)]
enum BinStore {
U16(Vec<u16>),
U32(Vec<u32>),
}
impl BinStore {
fn as_bins(&self) -> Bins<'_> {
match self {
BinStore::U16(v) => Bins::U16(v),
BinStore::U32(v) => Bins::U32(v),
}
}
#[inline]
fn get(&self, idx: usize) -> u32 {
match self {
BinStore::U16(v) => u32::from(v[idx]),
BinStore::U32(v) => v[idx],
}
}
fn len(&self) -> usize {
match self {
BinStore::U16(v) => v.len(),
BinStore::U32(v) => v.len(),
}
}
fn extend(&mut self, other: BinStore) {
match (self, other) {
(BinStore::U16(dst), BinStore::U16(src)) => dst.extend_from_slice(&src),
(BinStore::U32(dst), BinStore::U32(src)) => dst.extend_from_slice(&src),
_ => unreachable!("concatenated chunks share the index width"),
}
}
fn find(&self, s: usize, e: usize, fs: usize, fe: usize) -> Option<u32> {
match self {
BinStore::U16(v) => find_bin(&v[s..e], fs, fe),
BinStore::U32(v) => find_bin(&v[s..e], fs, fe),
}
}
}
#[inline]
fn find_bin<T: Copy + Into<u32>>(row: &[T], fs: usize, fe: usize) -> Option<u32> {
for &b in row {
let b: u32 = b.into();
if (b as usize) >= fs && (b as usize) < fe {
return Some(b);
}
}
None
}
pub enum Bins<'a> {
U16(&'a [u16]),
U32(&'a [u32]),
}
#[derive(Debug, Clone)]
pub struct GHistIndex {
n_rows: usize,
n_cols: usize,
row_ptr: Vec<usize>,
store: BinStore,
columns: Option<BinStore>,
cuts: HistCuts,
dense: bool,
}
impl GHistIndex {
pub fn from_dmatrix(data: &DMatrix, cuts: HistCuts) -> Self {
let n_rows = data.n_rows();
let n_cols = cuts.n_features();
let total_bins = cuts.total_bins();
let narrow = total_bins <= u16::MAX as usize + 1;
let threads = rayon::current_num_threads();
let search = BinSearch::new(&cuts);
let chunks: Vec<_> = if threads > 1 && n_rows.saturating_mul(n_cols) >= 65_536 {
let grain = n_rows.div_ceil(threads).max(1024);
(0..n_rows.div_ceil(grain))
.into_par_iter()
.map(|chunk| {
bin_rows(
data,
&search,
chunk * grain..((chunk + 1) * grain).min(n_rows),
narrow,
)
})
.collect()
} else {
vec![bin_rows(data, &search, 0..n_rows, narrow)]
};
drop(search);
let total = chunks.iter().map(|chunk| chunk.bins.len()).sum();
let dense = chunks.iter().all(|chunk| chunk.dense);
let max_bin = chunks.iter().map(|chunk| chunk.max_bin).max().unwrap_or(0);
assert!(
total == 0 || (max_bin as usize) < total_bins,
"binned index {max_bin} is outside the {total_bins} histogram bins"
);
let mut row_ptr = Vec::with_capacity(n_rows + 1);
row_ptr.push(0);
let mut offset = 0;
for chunk in &chunks {
row_ptr.extend(chunk.row_ends.iter().map(|end| offset + end));
offset += chunk.bins.len();
}
let mut store = if narrow {
BinStore::U16(Vec::with_capacity(total))
} else {
BinStore::U32(Vec::with_capacity(total))
};
for chunk in chunks {
store.extend(chunk.bins);
}
let columns = dense.then(|| match &store {
BinStore::U16(bins) => BinStore::U16(transpose_dense(bins, n_rows, n_cols)),
BinStore::U32(bins) => BinStore::U32(transpose_dense(bins, n_rows, n_cols)),
});
GHistIndex {
n_rows,
n_cols,
row_ptr,
store,
columns,
cuts,
dense,
}
}
#[inline]
pub fn n_rows(&self) -> usize {
self.n_rows
}
#[inline]
pub fn n_cols(&self) -> usize {
self.n_cols
}
#[inline]
pub fn cuts(&self) -> &HistCuts {
&self.cuts
}
#[inline]
pub fn total_bins(&self) -> usize {
self.cuts.total_bins()
}
#[inline]
pub fn row_ptr(&self) -> &[usize] {
&self.row_ptr
}
#[inline]
pub fn bins(&self) -> Bins<'_> {
self.store.as_bins()
}
#[inline]
pub fn dense_stride(&self) -> Option<usize> {
self.dense.then_some(self.n_cols)
}
#[inline]
pub fn column_bins(&self) -> Option<Bins<'_>> {
self.columns.as_ref().map(BinStore::as_bins)
}
#[inline]
pub fn row_len(&self, r: usize) -> usize {
self.row_ptr[r + 1] - self.row_ptr[r]
}
#[inline]
pub fn feature_bin_at(&self, r: usize, feature: usize, fs: usize, fe: usize) -> Option<u32> {
if self.dense {
return Some(self.store.get(self.row_ptr[r] + feature));
}
self.feature_bin(r, fs, fe)
}
#[inline]
pub fn feature_bin(&self, r: usize, fs: usize, fe: usize) -> Option<u32> {
let (s, e) = (self.row_ptr[r], self.row_ptr[r + 1]);
self.store.find(s, e, fs, fe)
}
}
pub(crate) fn transpose_dense<B: Copy + Default + Send + Sync>(
bins: &[B],
n_rows: usize,
n_cols: usize,
) -> Vec<B> {
const GROUP: usize = 8;
let mut columns = vec![B::default(); n_rows * n_cols];
if n_rows == 0 || n_cols == 0 {
return columns;
}
let fill = |(group, chunk): (usize, &mut [B])| {
let first = group * GROUP;
let width = chunk.len() / n_rows;
for r in 0..n_rows {
let row = &bins[r * n_cols + first..r * n_cols + first + width];
for (j, &bin) in row.iter().enumerate() {
chunk[j * n_rows + r] = bin;
}
}
};
if rayon::current_num_threads() > 1 && n_rows.saturating_mul(n_cols) >= 65_536 {
columns
.par_chunks_mut(GROUP * n_rows)
.enumerate()
.for_each(fill);
} else {
columns
.chunks_mut(GROUP * n_rows)
.enumerate()
.for_each(fill);
}
columns
}
struct BinnedRows {
row_ends: Vec<usize>,
bins: BinStore,
dense: bool,
max_bin: u32,
}
trait FromBin: Copy {
fn from_bin(bin: u32) -> Self;
}
macro_rules! impl_from_bin {
(narrow $t:ty) => {
impl FromBin for $t {
#[inline(always)]
fn from_bin(bin: u32) -> Self {
bin as $t
}
}
};
(wide $t:ty) => {
impl FromBin for $t {
#[inline(always)]
fn from_bin(bin: u32) -> Self {
bin
}
}
};
}
impl_from_bin!(narrow u16);
impl_from_bin!(wide u32);
fn bin_rows(data: &DMatrix, cuts: &BinSearch<'_>, rows: Range<usize>, narrow: bool) -> BinnedRows {
if narrow {
bin_rows_into::<u16>(data, cuts, rows, BinStore::U16)
} else {
bin_rows_into::<u32>(data, cuts, rows, BinStore::U32)
}
}
fn bin_rows_into<B: FromBin>(
data: &DMatrix,
cuts: &BinSearch<'_>,
rows: Range<usize>,
wrap: fn(Vec<B>) -> BinStore,
) -> BinnedRows {
let n_features = cuts.n_features();
let mut row_ends = Vec::with_capacity(rows.len());
let nnz = match data.csr_parts() {
Some((indptr, _, _)) => indptr[rows.end] - indptr[rows.start],
None => rows.len() * n_features,
};
let mut bins: Vec<B> = Vec::with_capacity(nnz);
let mut dense = true;
let mut max_bin = 0u32;
let mut push = |bins: &mut Vec<B>, bin: u32| {
max_bin = max_bin.max(bin);
bins.push(B::from_bin(bin));
};
if let Some(values) = data.dense_values() {
let missing = data.missing();
for r in rows {
let row = &values[r * n_features..(r + 1) * n_features];
let start = bins.len();
for (c, &v) in row.iter().enumerate() {
if !crate::data::dmatrix::is_missing(v, missing) {
push(&mut bins, cuts.bin_of(c, v));
}
}
dense &= bins.len() - start == n_features;
row_ends.push(bins.len());
}
} else {
let mut row: Vec<Entry> = Vec::new();
for r in rows {
data.row_into(r, &mut row);
dense &= row.len() == n_features
&& row.iter().enumerate().all(|(c, e)| e.index as usize == c);
for e in &row {
push(&mut bins, cuts.bin_of(e.index as usize, e.value));
}
row_ends.push(bins.len());
}
}
BinnedRows {
row_ends,
bins: wrap(bins),
dense,
max_bin,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parallel_binning_preserves_cuts_rows_and_width() {
use crate::data::FeatureType;
let serial = rayon::ThreadPoolBuilder::new()
.num_threads(1)
.build()
.unwrap();
let parallel = rayon::ThreadPoolBuilder::new()
.num_threads(4)
.build()
.unwrap();
for (rows, cols, missing, categorical) in [
(1025, 80, false, false),
(257, 300, false, false),
(1025, 80, true, false),
(1025, 80, false, true),
] {
let values: Vec<_> = (0..rows * cols)
.map(|i| {
if missing && i % 11 < 2 {
-1.0
} else if categorical && i % cols == 0 {
(i / cols % 4) as f32
} else {
((i / cols * 17 + i % cols * 31) % 509) as f32
}
})
.collect();
let mut data = DMatrix::from_dense_with_missing(&values, rows, cols, -1.0).unwrap();
if categorical {
let mut types = vec![FeatureType::Numerical; cols];
types[0] = FeatureType::Categorical;
data = data.with_feature_types(&types).unwrap();
}
let bin = || GHistIndex::from_dmatrix(&data, HistCuts::from_dmatrix(&data, 256));
let expected = serial.install(bin);
let actual = parallel.install(bin);
assert_eq!(
serde_json::to_value(actual.cuts()).unwrap(),
serde_json::to_value(expected.cuts()).unwrap()
);
assert_eq!(actual.row_ptr, expected.row_ptr);
assert_eq!(actual.dense, expected.dense);
match (&actual.store, &expected.store) {
(BinStore::U16(a), BinStore::U16(b)) => {
assert_eq!(a, b);
assert_eq!(cols, 80);
}
(BinStore::U32(a), BinStore::U32(b)) => {
assert_eq!(a, b);
assert_eq!(cols, 300);
}
_ => panic!("bin width changed"),
}
}
}
#[test]
fn bins_roundtrip_dense() {
let data = DMatrix::from_dense(&[0.0, 10.0, 1.0, 20.0], 2, 2).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
assert_eq!(ghist.row_len(0), 2);
assert_eq!(ghist.row_len(1), 2);
let b0 = ghist.cuts().bin_of(0, 0.0);
let b1 = ghist.cuts().bin_of(0, 1.0);
assert!(b1 > b0);
assert!(matches!(ghist.bins(), Bins::U16(_)));
}
#[test]
fn missing_entries_absent() {
let data = DMatrix::from_dense(&[0.0, f32::NAN, 1.0, 2.0], 2, 2).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
assert_eq!(ghist.row_len(0), 1);
assert_eq!(ghist.row_len(1), 2);
}
#[test]
fn feature_bin_lookup() {
let data = DMatrix::from_dense(&[0.0, 10.0, 1.0, 20.0], 2, 2).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 256);
let (f0s, f0e) = cuts.feature_bins(0);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
let b = ghist.feature_bin(0, f0s, f0e).unwrap();
assert!((b as usize) >= f0s && (b as usize) < f0e);
assert!(ghist.feature_bin(1, f0s, f0e).is_some());
}
}