use super::{BinIndex, feature_slices};
use crate::data::ghist::{Bins, GHistIndex};
use rayon::prelude::*;
use std::ops::Range;
pub(super) trait Bucket: Copy + Default + Send + Sync {
fn push(&mut self, value: Self);
}
pub(super) trait RowValue<E: Bucket>: Copy + Sync {
fn value(self) -> E;
}
const PREFETCH_ROWS: usize = 8;
const CACHE_LINE: usize = 64;
const TILE_ROWS: usize = 1024;
const BLOCK_BYTES: usize = 64 * 1024;
#[inline]
pub(super) fn accumulate<E: Bucket, V: RowValue<E>>(
ghist: &GHistIndex,
rows: &[u32],
values: &[V],
out: &mut [E],
) {
match ghist.bins() {
Bins::U16(bins) => accumulate_bins(ghist, bins, rows, values, out),
Bins::U32(bins) => accumulate_bins(ghist, bins, rows, values, out),
}
}
#[inline(always)]
fn prefetch_bins<B>(bins: &[B], 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);
}
}
}
fn feature_blocks(ghist: &GHistIndex, stride: usize, block_bins: usize) -> Vec<(usize, usize)> {
let cuts = ghist.cuts();
let mut blocks = 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));
blocks
}
#[inline(always)]
fn accumulate_bins<E: Bucket, V: RowValue<E>, B: BinIndex>(
ghist: &GHistIndex,
bins: &[B],
rows: &[u32],
values: &[V],
out: &mut [E],
) {
assert_eq!(
out.len(),
ghist.total_bins(),
"histogram length must equal the binned index's bin count"
);
let add_row = |row_bins: &[B], v: V, out: &mut [E]| {
let v = v.value();
let base = out.as_mut_ptr();
let (quads, rest) = row_bins.as_chunks::<4>();
for &quad in quads {
let [a, b, c, d] = quad.map(BinIndex::index);
unsafe {
let (ha, hb, hc, hd) = (*base.add(a), *base.add(b), *base.add(c), *base.add(d));
let add = |mut h: E| {
h.push(v);
h
};
*base.add(a) = add(ha);
*base.add(b) = add(hb);
*base.add(c) = add(hc);
*base.add(d) = add(hd);
}
}
for &bin in rest {
unsafe { &mut *base.add(bin.index()) }.push(v);
}
};
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, values, out),
Bins::U32(columns) => accumulate_columns(columns, n_rows, range, values, out),
}
return;
}
if let Some(stride) = ghist.dense_stride() {
accumulate_dense(ghist, bins, stride, rows, values, out, add_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_bins(bins, start, end - start);
}
if let Some(v) = values.get(ahead) {
crate::simd::prefetch_read(v);
}
}
let ri = r as usize;
add_row(&bins[rp[ri]..rp[ri + 1]], values[ri], out);
}
}
}
#[inline]
pub(super) fn contiguous_range(rows: &[u32]) -> Option<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<E: Bucket, V: RowValue<E>, B: BinIndex>(
columns: &[B],
n_rows: usize,
range: Range<usize>,
values: &[V],
out: &mut [E],
) {
fn column<'a, B>(group: &'a [B], k: usize, n_rows: usize, range: &Range<usize>) -> &'a [B] {
&group[k * n_rows..][..n_rows][range.clone()]
}
let values = &values[range.clone()];
let mut quads = columns.chunks_exact(4 * n_rows);
for quad in &mut quads {
let c = |k| column(quad, k, n_rows, &range);
sweep_columns([c(0), c(1), c(2), c(3)], values, out);
}
let rest = quads.remainder();
let c = |k| column(rest, k, n_rows, &range);
match rest.len() / n_rows {
3 => sweep_columns([c(0), c(1), c(2)], values, out),
2 => sweep_columns([c(0), c(1)], values, out),
1 => sweep_columns([c(0)], values, out),
_ => {}
}
}
#[inline(always)]
fn sweep_columns<const K: usize, E: Bucket, V: RowValue<E>, B: BinIndex>(
columns: [&[B]; K],
values: &[V],
out: &mut [E],
) {
let n = values.len();
let columns = columns.map(|c| &c[..n]);
let base = out.as_mut_ptr();
for (r, v) in values.iter().enumerate() {
let v = v.value();
let bins = columns.map(|c| unsafe { c.get_unchecked(r) }.index());
unsafe {
let mut entries = bins.map(|b| *base.add(b));
for e in &mut entries {
e.push(v);
}
for (&b, e) in bins.iter().zip(entries) {
*base.add(b) = e;
}
}
}
}
#[inline(always)]
fn accumulate_dense<E: Bucket, V: RowValue<E>, B: BinIndex>(
ghist: &GHistIndex,
bins: &[B],
stride: usize,
rows: &[u32],
values: &[V],
out: &mut [E],
add_row: impl Fn(&[B], V, &mut [E]),
) {
let blocks = feature_blocks(ghist, stride, BLOCK_BYTES / std::mem::size_of::<E>());
let short_rows = stride * std::mem::size_of::<B>() <= CACHE_LINE;
let prefetch = |rows: &[u32], i: usize| {
if let Some(&ahead) = rows.get(i + PREFETCH_ROWS) {
let ahead = ahead as usize;
if short_rows {
if let Some(row) = bins.get(ahead * stride..(ahead + 1) * stride)
&& let (Some(first), Some(last)) = (row.first(), row.last())
{
crate::simd::prefetch_read(first);
crate::simd::prefetch_read(last);
}
} else {
prefetch_bins(bins, ahead * stride, stride);
}
if let Some(v) = values.get(ahead) {
crate::simd::prefetch_read(v);
}
}
};
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], values[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], values[r as usize], out);
}
}
}
}
pub(super) enum SweepRows<'a> {
Range(Range<usize>),
Subset(&'a [u32]),
}
pub(super) fn by_features<E: Bucket, V: RowValue<E>>(
ghist: &GHistIndex,
columns: &Bins<'_>,
rows: &SweepRows<'_>,
values: &[V],
out: &mut [E],
threads: usize,
) {
let n_rows = ghist.n_rows();
let per_task = ghist.n_cols().div_ceil(threads).clamp(1, 4);
let mut slices = feature_slices(ghist, out, 1);
slices
.par_chunks_mut(per_task)
.enumerate()
.for_each(|(task, features)| {
let f = task * per_task;
match columns {
Bins::U16(c) => feature_group(c, n_rows, f, features, rows, values),
Bins::U32(c) => feature_group(c, n_rows, f, features, rows, values),
}
});
}
#[inline(always)]
fn feature_group<E: Bucket, V: RowValue<E>, B: BinIndex>(
columns: &[B],
n_rows: usize,
f: usize,
features: &mut [(usize, &mut [E])],
rows: &SweepRows<'_>,
values: &[V],
) {
let column = |f: usize| &columns[f * n_rows..][..n_rows];
for (_, slice) in features.iter_mut() {
slice.fill(E::default());
}
match features {
[a] => sweep_slices([column(f)], [a], rows, values),
[a, b] => sweep_slices([column(f), column(f + 1)], [a, b], rows, values),
[a, b, c] => sweep_slices(
[column(f), column(f + 1), column(f + 2)],
[a, b, c],
rows,
values,
),
[a, b, c, d] => sweep_slices(
[column(f), column(f + 1), column(f + 2), column(f + 3)],
[a, b, c, d],
rows,
values,
),
_ => unreachable!("callers pass one to four features"),
}
}
#[inline(always)]
fn sweep_slices<const K: usize, E: Bucket, V: RowValue<E>, B: BinIndex>(
columns: [&[B]; K],
slices: [&mut (usize, &mut [E]); K],
rows: &SweepRows<'_>,
values: &[V],
) {
let mut slices = slices.map(|(first, slice)| (*first, &mut **slice));
match rows {
SweepRows::Range(range) => {
let values = &values[range.clone()];
let n = values.len();
let columns = columns.map(|c| &c[range.clone()][..n]);
for (r, v) in values.iter().enumerate() {
let v = v.value();
for (column, (first, slice)) in columns.iter().zip(&mut slices) {
slice[column[r].index() - *first].push(v);
}
}
}
SweepRows::Subset(rows) => {
for &r in *rows {
let r = r as usize;
let v = values[r].value();
for (column, (first, slice)) in columns.iter().zip(&mut slices) {
slice[column[r].index() - *first].push(v);
}
}
}
}
}