use crate::data::ghist::{Bins, GHistIndex};
use crate::objective::GradPair;
use crate::tree::gain::GradStats;
use rayon::prelude::*;
pub type Histogram = Vec<GradStats>;
pub fn zeroed(total_bins: usize) -> Histogram {
vec![GradStats::default(); total_bins]
}
pub fn subtract_in_place(parent: &mut [GradStats], child: &[GradStats]) {
debug_assert_eq!(parent.len(), child.len());
for (p, c) in parent.iter_mut().zip(child) {
*p = p.sub(*c);
}
}
pub trait HistogramBackend: Send + Sync {
fn build(&self, ghist: &GHistIndex, rows: &[u32], gpair: &[GradPair], out: &mut [GradStats]);
}
#[derive(Debug, Clone, Copy, Default)]
pub struct CpuBackend;
const ROWS_PER_TASK: usize = 4096;
const PARALLEL_THRESHOLD: usize = 2 * ROWS_PER_TASK;
const REDUCE_BINS: usize = 2048;
impl HistogramBackend for CpuBackend {
fn build(&self, ghist: &GHistIndex, rows: &[u32], gpair: &[GradPair], out: &mut [GradStats]) {
let total = out.len();
let threads = rayon::current_num_threads();
if threads <= 1 || rows.len() < PARALLEL_THRESHOLD {
out.fill(GradStats::default());
accumulate(ghist, rows, gpair, out);
return;
}
if let (Some(columns), Some(range)) = (ghist.column_bins(), contiguous_range(rows)) {
let n_rows = ghist.n_rows();
let cuts = ghist.cuts();
let mut slices = Vec::with_capacity(ghist.n_cols());
let mut rest = out;
let mut next = 0;
for f in 0..ghist.n_cols() {
let (fs, fe) = cuts.feature_bins(f);
assert_eq!(
fs, next,
"feature bin ranges must be contiguous and ordered"
);
let (head, tail) = rest.split_at_mut(fe - fs);
slices.push((fs, head));
rest = tail;
next = fe;
}
assert!(
rest.is_empty(),
"feature bin ranges must cover the histogram"
);
slices
.into_par_iter()
.enumerate()
.for_each(|(f, (fs, slice))| {
slice.fill(GradStats::default());
match columns {
Bins::U16(c) => {
accumulate_column(&c[f * n_rows..][..n_rows], fs, &range, gpair, slice);
}
Bins::U32(c) => {
accumulate_column(&c[f * n_rows..][..n_rows], fs, &range, gpair, slice);
}
}
});
return;
}
let tasks = threads.min(rows.len() / ROWS_PER_TASK);
let grain = rows.len().div_ceil(tasks);
let partials: Vec<Histogram> = rows
.par_chunks(grain)
.map(|chunk| {
let mut local = zeroed(total);
accumulate(ghist, chunk, gpair, &mut local);
local
})
.collect();
out.par_chunks_mut(REDUCE_BINS)
.enumerate()
.for_each(|(i, out)| {
let start = i * REDUCE_BINS;
let end = start + out.len();
out.copy_from_slice(&partials[0][start..end]);
for partial in &partials[1..] {
for (o, p) in out.iter_mut().zip(&partial[start..end]) {
o.add(*p);
}
}
});
}
}
const PREFETCH_ROWS: usize = 8;
const CACHE_LINE: usize = 64;
pub(crate) trait BinIndex: Copy + Send + Sync {
fn index(self) -> usize;
}
impl BinIndex for u16 {
#[inline(always)]
fn index(self) -> usize {
self as usize
}
}
impl BinIndex for u32 {
#[inline(always)]
fn index(self) -> usize {
self as usize
}
}
#[inline]
fn accumulate(ghist: &GHistIndex, rows: &[u32], gpair: &[GradPair], out: &mut [GradStats]) {
match ghist.bins() {
Bins::U16(bins) => accumulate_bins(ghist, bins, rows, gpair, out),
Bins::U32(bins) => accumulate_bins(ghist, bins, rows, gpair, out),
}
}
#[inline(always)]
fn accumulate_bins<B: BinIndex>(
ghist: &GHistIndex,
bins: &[B],
rows: &[u32],
gpair: &[GradPair],
out: &mut [GradStats],
) {
assert_eq!(
out.len(),
ghist.total_bins(),
"histogram length must equal the binned index's bin count"
);
let add_row = |row_bins: &[B], gp: GradPair, out: &mut [GradStats]| {
let g = GradStats::from_pair(gp);
for &bin in row_bins {
unsafe { out.get_unchecked_mut(bin.index()) }.add(g);
}
};
let prefetch_row = |start: usize, len: usize| {
for offset in (0..len).step_by(CACHE_LINE / std::mem::size_of::<B>()) {
if let Some(bin) = bins.get(start + offset) {
crate::simd::prefetch_read(bin);
}
}
};
if let Some(columns) = ghist.column_bins()
&& let Some(range) = contiguous_range(rows)
{
let n_rows = ghist.n_rows();
match columns {
Bins::U16(columns) => accumulate_columns(columns, n_rows, range, gpair, out),
Bins::U32(columns) => accumulate_columns(columns, n_rows, range, gpair, out),
}
return;
}
if let Some(stride) = ghist.dense_stride() {
accumulate_dense(ghist, bins, stride, rows, gpair, out, add_row, prefetch_row);
} else {
let rp = ghist.row_ptr();
for (i, &r) in rows.iter().enumerate() {
if let Some(&ahead) = rows.get(i + PREFETCH_ROWS) {
let ahead = ahead as usize;
if let (Some(&start), Some(&end)) = (rp.get(ahead), rp.get(ahead + 1)) {
prefetch_row(start, end - start);
}
if let Some(gp) = gpair.get(ahead) {
crate::simd::prefetch_read(gp);
}
}
let ri = r as usize;
add_row(&bins[rp[ri]..rp[ri + 1]], gpair[ri], out);
}
}
}
#[inline]
fn contiguous_range(rows: &[u32]) -> Option<std::ops::Range<usize>> {
let first = *rows.first()? as usize;
let end = first.checked_add(rows.len())?;
let contiguous = rows
.iter()
.enumerate()
.all(|(i, &row)| row as usize == first + i);
contiguous.then_some(first..end)
}
#[inline(always)]
fn accumulate_columns<B: BinIndex>(
columns: &[B],
n_rows: usize,
range: std::ops::Range<usize>,
gpair: &[GradPair],
out: &mut [GradStats],
) {
let gpair = &gpair[range.clone()];
for column in columns.chunks_exact(n_rows) {
for (&bin, gp) in column[range.clone()].iter().zip(gpair) {
let g = GradStats::from_pair(*gp);
unsafe { out.get_unchecked_mut(bin.index()) }.add(g);
}
}
}
#[inline(always)]
fn accumulate_column<B: BinIndex>(
column: &[B],
first_bin: usize,
range: &std::ops::Range<usize>,
gpair: &[GradPair],
slice: &mut [GradStats],
) {
for (&bin, gp) in column[range.clone()].iter().zip(&gpair[range.clone()]) {
let g = GradStats::from_pair(*gp);
slice[bin.index() - first_bin].add(g);
}
}
const TILE_ROWS: usize = 4096;
const BLOCK_BINS: usize = 4096;
#[inline(always)]
#[allow(clippy::too_many_arguments)]
fn accumulate_dense<B: BinIndex>(
ghist: &GHistIndex,
bins: &[B],
stride: usize,
rows: &[u32],
gpair: &[GradPair],
out: &mut [GradStats],
add_row: impl Fn(&[B], GradPair, &mut [GradStats]),
prefetch_row: impl Fn(usize, usize),
) {
let cuts = ghist.cuts();
let mut blocks: Vec<(usize, usize)> = Vec::new();
let mut block_start = 0;
for f in 1..=stride {
let span = cuts.feature_bins(f - 1).1 - cuts.feature_bins(block_start).0;
if span > BLOCK_BINS && f - 1 > block_start {
blocks.push((block_start, f - 1));
block_start = f - 1;
}
}
blocks.push((block_start, stride));
let prefetch = |rows: &[u32], i: usize| {
if let Some(&ahead) = rows.get(i + PREFETCH_ROWS) {
let ahead = ahead as usize;
prefetch_row(ahead * stride, stride);
if let Some(gp) = gpair.get(ahead) {
crate::simd::prefetch_read(gp);
}
}
};
if blocks.len() == 1 {
for (i, &r) in rows.iter().enumerate() {
prefetch(rows, i);
let start = r as usize * stride;
add_row(&bins[start..start + stride], gpair[r as usize], out);
}
return;
}
for tile in rows.chunks(TILE_ROWS) {
for (block, &(f0, f1)) in blocks.iter().enumerate() {
for (i, &r) in tile.iter().enumerate() {
if block == 0 {
prefetch(tile, i);
}
let start = r as usize * stride;
add_row(&bins[start + f0..start + f1], gpair[r as usize], out);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::data::DMatrix;
use crate::data::quantile::HistCuts;
fn brute_force(
ghist: &GHistIndex,
rows: &[u32],
gpair: &[GradPair],
total: usize,
) -> Histogram {
let mut h = zeroed(total);
accumulate(ghist, rows, gpair, &mut h);
h
}
#[test]
fn build_matches_brute_force() {
let n = 200;
let x: Vec<f32> = (0..n).map(|i| (i % 17) as f32).collect();
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 32);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
let gpair: Vec<GradPair> = (0..n)
.map(|i| GradPair::new((i as f32) * 0.1 - 5.0, 1.0))
.collect();
let rows: Vec<u32> = (0..n as u32).collect();
let mut out = zeroed(ghist.total_bins());
CpuBackend.build(&ghist, &rows, &gpair, &mut out);
let expect = brute_force(&ghist, &rows, &gpair, ghist.total_bins());
for (a, b) in out.iter().zip(&expect) {
assert!((a.grad - b.grad).abs() < 1e-4);
assert!((a.hess - b.hess).abs() < 1e-4);
}
}
#[test]
fn subtraction_identity() {
let total = 8;
let mut parent = zeroed(total);
let mut left = zeroed(total);
let mut right = zeroed(total);
for i in 0..total {
left[i] = GradStats::new(i as f64, 1.0);
right[i] = GradStats::new(-(i as f64) * 0.5, 2.0);
parent[i] = GradStats::new(left[i].grad + right[i].grad, left[i].hess + right[i].hess);
}
let mut out = parent.clone();
subtract_in_place(&mut out, &left);
for i in 0..total {
assert!((out[i].grad - right[i].grad).abs() < 1e-12);
assert!((out[i].hess - right[i].hess).abs() < 1e-12);
}
}
#[test]
fn parallel_matches_sequential_large() {
let n = PARALLEL_THRESHOLD + 37;
let x: Vec<f32> = (0..n).map(|i| (i % 251) as f32).collect();
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 64);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
let gpair: Vec<GradPair> = (0..n)
.map(|i| GradPair::new(((i * 7) % 13) as f32 - 6.0, 1.0))
.collect();
let rows: Vec<u32> = (0..n as u32).collect();
let mut out = zeroed(ghist.total_bins());
rayon::ThreadPoolBuilder::new()
.num_threads(4)
.build()
.unwrap()
.install(|| CpuBackend.build(&ghist, &rows, &gpair, &mut out));
let expect = brute_force(&ghist, &rows, &gpair, ghist.total_bins());
for (a, b) in out.iter().zip(&expect) {
assert!(
(a.grad - b.grad).abs() < 1e-2,
"grad {} vs {}",
a.grad,
b.grad
);
assert!((a.hess - b.hess).abs() < 1e-2);
}
}
fn row_order_reference(ghist: &GHistIndex, rows: &[u32], gpair: &[GradPair]) -> Histogram {
let stride = ghist.dense_stride().expect("dense index");
let mut h = zeroed(ghist.total_bins());
for &r in rows {
let r = r as usize;
let g = GradStats::from_pair(gpair[r]);
for f in 0..stride {
let bin = match ghist.bins() {
Bins::U16(b) => b[r * stride + f] as usize,
Bins::U32(b) => b[r * stride + f] as usize,
};
h[bin].add(g);
}
}
h
}
#[test]
fn column_and_row_sweeps_match_reference_bit_for_bit() {
let (n, f) = (3 * PARALLEL_THRESHOLD + 129, 7);
let x: Vec<f32> = (0..n * f)
.map(|i| ((i * 2_654_435_761_usize) % 1009) as f32 / 7.0)
.collect();
let data = DMatrix::from_dense(&x, n, f).unwrap();
let cuts = HistCuts::from_dmatrix(&data, 64);
let ghist = GHistIndex::from_dmatrix(&data, cuts);
assert!(
ghist.column_bins().is_some(),
"dense index keeps a column copy"
);
let gpair: Vec<GradPair> = (0..n)
.map(|i| {
GradPair::new(
((i * 7919) % 1237) as f32 / 331.0 - 1.9,
0.25 + (i % 5) as f32,
)
})
.collect();
let bits = |h: &Histogram| -> Vec<(u64, u64)> {
h.iter()
.map(|s| (s.grad.to_bits(), s.hess.to_bits()))
.collect()
};
let all: Vec<u32> = (0..n as u32).collect();
let subset: Vec<u32> = (0..n as u32).filter(|r| r % 3 != 1).collect();
let offset: Vec<u32> = (1000..(1000 + PARALLEL_THRESHOLD) as u32).collect();
assert!(contiguous_range(&all).is_some() && contiguous_range(&offset).is_some());
assert!(contiguous_range(&subset).is_none());
assert!(contiguous_range(&[2, 0, 1]).is_none() && contiguous_range(&[5, 5]).is_none());
for rows in [&all, &subset, &offset] {
let expect = bits(&row_order_reference(&ghist, rows, &gpair));
for threads in [1, 4] {
let mut out = zeroed(ghist.total_bins());
rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap()
.install(|| CpuBackend.build(&ghist, rows, &gpair, &mut out));
assert_eq!(bits(&out), expect, "rows={} threads={threads}", rows.len());
}
}
}
}