use super::walk::{Bucket, RowValue, accumulate, contiguous_range};
use super::{BinIndex, PARALLEL_THRESHOLD, REDUCE_BINS, ROWS_PER_TASK, feature_slices};
use crate::config::QuantizedGrad;
use crate::data::ghist::{Bins, GHistIndex};
use crate::objective::GradPair;
use crate::rng::{keyed_unit, mix64};
use crate::tree::gain::GradStats;
use rayon::prelude::*;
use std::sync::Arc;
const QUANTIZE_CHUNK: usize = 8192;
const GRAD_STREAM: u64 = 0x6772_6164_5F71_6E74;
const HESS_STREAM: u64 = 0x6865_7373_5F71_6E74;
#[inline]
fn round_signed(x: f64, u: f64) -> i32 {
if x >= 0.0 {
(x + u) as i32
} else {
(x - u) as i32
}
}
#[inline(always)]
fn row_parts(p: i32) -> (i64, i64) {
let h = i64::from(p as i16);
(i64::from((p - h as i32) >> 16), h)
}
#[derive(Debug)]
pub(crate) struct QuantizedGradients {
packed: Vec<i32>,
grad_scale: f64,
hess_scale: f64,
row_bound: u64,
}
impl QuantizedGradients {
fn quantize(gpair: &[GradPair], bins: usize, stochastic: bool, seed: u64) -> Self {
debug_assert!((2..=127).contains(&bins), "validated by TrainingParams");
let first_hess = gpair.first().map_or(0.0, |gp| gp.hess);
let (max_grad, max_hess, constant) = gpair
.par_chunks(QUANTIZE_CHUNK)
.map(|chunk| {
chunk.iter().fold((0f32, 0f32, true), |(g, h, c), gp| {
(
g.max(gp.grad.abs()),
h.max(gp.hess.abs()),
c && gp.hess == first_hess,
)
})
})
.reduce(
|| (0.0, 0.0, true),
|(g1, h1, c1), (g2, h2, c2)| (g1.max(g2), h1.max(h2), c1 && c2),
);
let half = (bins / 2) as i32;
let levels = bins as i32;
let scale_of = |max: f32, divisor: i32| {
if max > 0.0 {
(max / divisor as f32).max(f32::from_bits(1))
} else {
0.0
}
};
let grad_scale = scale_of(max_grad, half);
let hess_scale = if constant {
first_hess
} else {
scale_of(max_hess, levels)
};
let inverse = |scale: f32| {
if scale > 0.0 && scale.is_finite() {
1.0 / f64::from(scale)
} else {
0.0
}
};
let (inv_grad, inv_hess) = (inverse(grad_scale), inverse(hess_scale));
let (grad_key, hess_key) = (mix64(seed ^ GRAD_STREAM), mix64(seed ^ HESS_STREAM));
let mut packed = vec![0i32; gpair.len()];
packed
.par_chunks_mut(QUANTIZE_CHUNK)
.zip(gpair.par_chunks(QUANTIZE_CHUNK))
.enumerate()
.for_each(|(chunk, (out, grads))| {
let first = chunk * QUANTIZE_CHUNK;
for (i, (o, gp)) in out.iter_mut().zip(grads).enumerate() {
let row = first + i;
let (ug, uh) = if stochastic {
(
keyed_unit(grad_key, row as u64),
keyed_unit(hess_key, row as u64),
)
} else {
(0.5, 0.5)
};
let g = round_signed(f64::from(gp.grad) * inv_grad, ug).clamp(-half, half);
let h = if constant {
1
} else {
round_signed(f64::from(gp.hess) * inv_hess, uh).clamp(-levels, levels)
};
*o = (g << 16) + h;
}
});
QuantizedGradients {
packed,
grad_scale: f64::from(grad_scale),
hess_scale: f64::from(hess_scale),
row_bound: if constant { half.max(1) } else { levels } as u64,
}
}
fn width(&self, rows: usize) -> Width {
let bound = rows as u64 * self.row_bound;
if bound < 1 << 15 {
Width::W32
} else if bound < 1 << 31 {
Width::W64
} else {
Width::W128
}
}
#[cfg(test)]
fn row_stats(&self, row: usize) -> GradStats {
let (g, h) = row_parts(self.packed[row]);
self.dequantize_parts(g, h)
}
#[inline]
fn dequantize_parts(&self, g: i64, h: i64) -> GradStats {
GradStats::new(g as f64 * self.grad_scale, h as f64 * self.hess_scale)
}
fn node_stats(&self, rows: &[u32]) -> GradStats {
let (g, h) = rows.iter().fold((0i64, 0i64), |(g, h), &r| {
let (rg, rh) = row_parts(self.packed[r as usize]);
(g + rg, h + rh)
});
self.dequantize_parts(g, h)
}
fn dequantize(&self, hist: &QuantHist) -> Vec<GradStats> {
let mut out = Vec::new();
self.dequantize_into(hist, &mut out);
out
}
fn dequantize_into(&self, hist: &QuantHist, out: &mut Vec<GradStats>) {
fn map<A: Packed>(q: &QuantizedGradients, bins: &[A], out: &mut Vec<GradStats>) {
out.clear();
out.extend(bins.iter().map(|b| {
let (g, h) = b.parts();
q.dequantize_parts(g, h)
}));
}
match hist {
QuantHist::W32(b) => map(self, b, out),
QuantHist::W64(b) => map(self, b, out),
QuantHist::W128(b) => map(self, b, out),
}
}
fn build(&self, ghist: &GHistIndex, rows: &[u32]) -> QuantHist {
match self.width(rows.len()) {
Width::W32 => QuantHist::W32(self.build_typed(ghist, rows)),
Width::W64 => QuantHist::W64(self.build_typed(ghist, rows)),
Width::W128 => QuantHist::W128(self.build_typed(ghist, rows)),
}
}
fn build_typed<A: Packed>(&self, ghist: &GHistIndex, rows: &[u32]) -> Vec<A> {
let total = ghist.total_bins();
let threads = rayon::current_num_threads();
if threads <= 1 || rows.len() < PARALLEL_THRESHOLD {
let mut out = vec![A::default(); total];
self.accumulate_narrow(ghist, rows, &mut out);
return out;
}
if let (Some(columns), Some(range)) = (ghist.column_bins(), contiguous_range(rows)) {
let mut out = vec![A::default(); total];
let values: Vec<A> = self.packed[range.clone()]
.iter()
.map(|&p| A::from_row(p))
.collect();
let n_rows = ghist.n_rows();
feature_slices(ghist, &mut out, 1)
.into_par_iter()
.enumerate()
.for_each(|(f, (fs, slice))| {
let span = range.clone();
match columns {
Bins::U16(c) => {
accumulate_column(&c[f * n_rows..][span], fs, &values, slice);
}
Bins::U32(c) => {
accumulate_column(&c[f * n_rows..][span], fs, &values, slice);
}
}
});
return out;
}
let tasks = threads.min(rows.len() / ROWS_PER_TASK);
let grain = rows.len().div_ceil(tasks);
match self.width(grain) {
Width::W32 => self.reduce_chunks::<A, i32>(ghist, rows, grain),
Width::W64 => self.reduce_chunks::<A, i64>(ghist, rows, grain),
Width::W128 => self.reduce_chunks::<A, i128>(ghist, rows, grain),
}
}
fn reduce_chunks<A: Packed, P: Packed>(
&self,
ghist: &GHistIndex,
rows: &[u32],
grain: usize,
) -> Vec<A> {
let total = ghist.total_bins();
let partials: Vec<Vec<P>> = rows
.par_chunks(grain)
.map(|chunk| {
let mut local = vec![P::default(); total];
self.accumulate_narrow(ghist, chunk, &mut local);
local
})
.collect();
let mut out = vec![A::default(); total];
out.par_chunks_mut(REDUCE_BINS)
.enumerate()
.for_each(|(i, out)| {
let start = i * REDUCE_BINS;
for partial in &partials {
for (o, &p) in out.iter_mut().zip(&partial[start..]) {
*o = o.add(A::repack(p));
}
}
});
out
}
fn accumulate_narrow<A: Packed>(&self, ghist: &GHistIndex, rows: &[u32], out: &mut [A]) {
let run = ((1 << 15) - 1) / self.row_bound as usize;
if A::BITS == 32 || rows.len() <= run {
accumulate(ghist, rows, &self.packed, out);
return;
}
let mut scratch = vec![0i32; out.len()];
for chunk in rows.chunks(run) {
accumulate(ghist, chunk, &self.packed, &mut scratch);
for (o, s) in out.iter_mut().zip(&mut scratch) {
*o = o.add(A::repack(*s));
*s = 0;
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
enum Width {
W32,
W64,
W128,
}
#[derive(Debug, PartialEq)]
pub(crate) enum QuantHist {
W32(Vec<i32>),
W64(Vec<i64>),
W128(Vec<i128>),
}
impl QuantHist {
fn width(&self) -> Width {
match self {
QuantHist::W32(_) => Width::W32,
QuantHist::W64(_) => Width::W64,
QuantHist::W128(_) => Width::W128,
}
}
fn widen_to(&mut self, width: Width) {
fn convert<A: Packed, P: Packed>(bins: &[P]) -> Vec<A> {
bins.iter().map(|&p| A::repack(p)).collect()
}
if self.width() >= width {
return;
}
*self = match (&*self, width) {
(QuantHist::W32(b), Width::W64) => QuantHist::W64(convert(b)),
(QuantHist::W32(b), _) => QuantHist::W128(convert(b)),
(QuantHist::W64(b), _) => QuantHist::W128(convert(b)),
(QuantHist::W128(_), _) => unreachable!("W128 is the widest width"),
};
}
fn subtract(&mut self, child: &QuantHist) {
fn sub<A: Packed>(parent: &mut [A], child: &QuantHist) {
fn typed<A: Packed, P: Packed>(parent: &mut [A], child: &[P]) {
debug_assert_eq!(parent.len(), child.len());
for (p, &c) in parent.iter_mut().zip(child) {
*p = p.sub(A::repack(c));
}
}
match child {
QuantHist::W32(c) => typed(parent, c),
QuantHist::W64(c) => typed(parent, c),
QuantHist::W128(c) => typed(parent, c),
}
}
self.widen_to(child.width());
match self {
QuantHist::W32(p) => sub(p, child),
QuantHist::W64(p) => sub(p, child),
QuantHist::W128(p) => sub(p, child),
}
}
}
pub(crate) type QuantChild = (QuantNode, Vec<GradStats>);
#[derive(Debug)]
pub(crate) struct QuantNode {
grads: Arc<QuantizedGradients>,
hist: QuantHist,
}
impl QuantNode {
pub(crate) fn root(
ghist: &GHistIndex,
gpair: &[GradPair],
rows: &[u32],
quantized: QuantizedGrad,
seed: u64,
) -> (Self, GradStats, Vec<GradStats>) {
let grads = Arc::new(QuantizedGradients::quantize(
gpair,
quantized.bins(),
quantized.stochastic_rounding(),
seed,
));
let hist = grads.build(ghist, rows);
let stats = grads.node_stats(rows);
let float = grads.dequantize(&hist);
(QuantNode { grads, hist }, stats, float)
}
pub(crate) fn children(
self,
ghist: &GHistIndex,
left_rows: &[u32],
right_rows: &[u32],
spare: Vec<GradStats>,
) -> (QuantChild, QuantChild) {
let QuantNode {
grads,
hist: mut sibling,
} = self;
let left_smaller = left_rows.len() <= right_rows.len();
let small = grads.build(ghist, if left_smaller { left_rows } else { right_rows });
sibling.subtract(&small);
let small_float = grads.dequantize(&small);
let mut sibling_float = spare;
grads.dequantize_into(&sibling, &mut sibling_float);
let small = QuantNode {
grads: Arc::clone(&grads),
hist: small,
};
let sibling = QuantNode {
grads,
hist: sibling,
};
if left_smaller {
((small, small_float), (sibling, sibling_float))
} else {
((sibling, sibling_float), (small, small_float))
}
}
}
trait Packed: Bucket {
const BITS: u32;
fn from_parts(g: i64, h: i64) -> Self;
fn parts(self) -> (i64, i64);
fn from_row(p: i32) -> Self;
fn add(self, other: Self) -> Self;
fn sub(self, other: Self) -> Self;
#[inline]
fn repack<P: Packed>(p: P) -> Self {
let (g, h) = p.parts();
Self::from_parts(g, h)
}
}
macro_rules! packed {
($ty:ty, $half:ty, $shift:expr) => {
impl Packed for $ty {
const BITS: u32 = <$ty>::BITS;
#[inline(always)]
fn from_parts(g: i64, h: i64) -> Self {
((g as $ty) << $shift) + h as $ty
}
#[inline(always)]
#[allow(clippy::cast_lossless, reason = "the same cast narrows for i128")]
fn parts(self) -> (i64, i64) {
let h = self as $half;
(((self - <$ty>::from(h)) >> $shift) as i64, i64::from(h))
}
#[inline(always)]
fn from_row(p: i32) -> Self {
let (g, h) = row_parts(p);
Self::from_parts(g, h)
}
#[inline(always)]
fn add(self, other: Self) -> Self {
self + other
}
#[inline(always)]
fn sub(self, other: Self) -> Self {
self - other
}
}
impl Bucket for $ty {
#[inline(always)]
fn push(&mut self, value: Self) {
*self += value;
}
}
};
}
packed!(i32, i16, 16);
packed!(i64, i32, 32);
packed!(i128, i64, 64);
#[inline(always)]
fn accumulate_column<A: Packed, B: BinIndex>(
column: &[B],
first_bin: usize,
values: &[A],
slice: &mut [A],
) {
for (&bin, &v) in column.iter().zip(values) {
let slot = &mut slice[bin.index() - first_bin];
*slot = slot.add(v);
}
}
impl<A: Packed> RowValue<A> for i32 {
#[inline(always)]
fn value(self) -> A {
A::from_row(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::data::DMatrix;
use crate::data::quantile::HistCuts;
fn gp(grad: f32, hess: f32) -> GradPair {
GradPair::new(grad, hess)
}
fn gradients(n: usize, seed: u64) -> Vec<GradPair> {
(0..n)
.map(|i| {
let a = keyed_unit(seed, i as u64) as f32;
let b = keyed_unit(seed ^ 1, i as u64) as f32;
gp(4.0 * a - 1.5, 0.05 + b)
})
.collect()
}
fn binned(n: usize, features: usize, missing: bool) -> GHistIndex {
let x: Vec<f32> = (0..n * features)
.map(|i| {
if missing && i % 7 == 3 {
f32::NAN
} else {
keyed_unit(99, i as u64) as f32
}
})
.collect();
let data = DMatrix::from_dense(&x, n, features).unwrap();
GHistIndex::from_dmatrix(&data, HistCuts::from_dmatrix(&data, 64))
}
#[test]
fn packed_widths_round_trip_signed_halves() {
for (g, h) in [(0, 0), (-3, 5), (7, -2), (-1, -1), (12_000, 16_000)] {
assert_eq!(i32::from_parts(g, h).parts(), (g, h));
assert_eq!(
i64::from_parts(g << 14, h << 14).parts(),
(g << 14, h << 14)
);
assert_eq!(
i128::from_parts(g << 40, h << 40).parts(),
(g << 40, h << 40)
);
}
let a = i32::from_parts(-5, 3);
let b = i32::from_parts(2, -7);
assert_eq!(a.add(b).parts(), (-3, -4));
assert_eq!(a.sub(b).parts(), (-7, 10));
assert_eq!(i64::repack(a).parts(), (-5, 3));
}
#[test]
fn deterministic_rounding_error_is_at_most_half_a_step() {
let gpair = gradients(10_000, 3);
for bins in [2, 4, 16, 127] {
let q = QuantizedGradients::quantize(&gpair, bins, false, 0);
let half = (bins / 2) as f64;
for (row, g) in gpair.iter().enumerate() {
let s = q.row_stats(row);
assert!((s.grad - f64::from(g.grad)).abs() <= 0.5 * q.grad_scale * (1.0 + 1e-6));
assert!((s.hess - f64::from(g.hess)).abs() <= 0.5 * q.hess_scale * (1.0 + 1e-6));
let (qg, qh) = row_parts(q.packed[row]);
assert!(qg.abs() as f64 <= half && qh >= 0 && qh <= bins as i64);
}
}
}
#[test]
fn stochastic_rounding_is_unbiased_and_within_one_step() {
let gpair = vec![gp(1.0, 1.0), gp(0.37, 0.61), gp(-0.83, 0.2)];
let trials = 20_000;
let mut sums = [(0f64, 0f64); 3];
for seed in 0..trials {
let q = QuantizedGradients::quantize(&gpair, 4, true, seed);
for (row, sum) in sums.iter_mut().enumerate() {
let s = q.row_stats(row);
assert!((s.grad - f64::from(gpair[row].grad)).abs() < q.grad_scale);
assert!((s.hess - f64::from(gpair[row].hess)).abs() < q.hess_scale);
sum.0 += s.grad;
sum.1 += s.hess;
}
}
let n = trials as f64;
for (row, (g, h)) in sums.iter().enumerate() {
assert!(
(g / n - f64::from(gpair[row].grad)).abs() < 0.01,
"row {row}: {}",
g / n
);
assert!(
(h / n - f64::from(gpair[row].hess)).abs() < 0.01,
"row {row}: {}",
h / n
);
}
let q = QuantizedGradients::quantize(&[gp(1.0, 1.0), gp(-0.5, 0.25)], 4, true, 7);
assert_eq!(q.row_stats(0), GradStats::new(1.0, 1.0));
assert_eq!(q.row_stats(1), GradStats::new(-0.5, 0.25));
}
#[test]
fn constant_hessians_are_exact() {
let gpair: Vec<GradPair> = (0..100).map(|i| gp(i as f32 - 50.0, 0.75)).collect();
let q = QuantizedGradients::quantize(&gpair, 4, true, 1);
for row in 0..gpair.len() {
assert_eq!(q.row_stats(row).hess, 0.75);
}
let zeros = QuantizedGradients::quantize(&[gp(0.0, 0.0); 5], 4, true, 1);
assert_eq!(zeros.node_stats(&[0, 1, 2, 3, 4]), GradStats::default());
}
#[test]
fn negative_hessians_keep_their_sign() {
let gpair = vec![gp(0.5, -1.0), gp(-0.5, 2.0), gp(0.25, -0.5)];
let q = QuantizedGradients::quantize(&gpair, 8, false, 0);
assert_eq!(q.node_stats(&[0, 1, 2]), GradStats::new(0.25, 0.5));
}
#[test]
fn widths_follow_row_count_and_levels() {
let gpair = gradients(8, 5);
let q = QuantizedGradients::quantize(&gpair, 4, true, 0);
assert_eq!(q.width(8191), Width::W32);
assert_eq!(q.width(8192), Width::W64);
assert_eq!(q.width((1 << 29) - 1), Width::W64);
assert_eq!(q.width(1 << 29), Width::W128);
let constant = QuantizedGradients::quantize(&[gp(1.0, 1.0); 4], 4, true, 0);
assert_eq!(constant.width(16_383), Width::W32);
assert_eq!(constant.width(16_384), Width::W64);
}
fn assert_matches_reference(ghist: &GHistIndex, q: &QuantizedGradients, rows: &[u32]) {
let cuts = ghist.cuts();
let mut expected = vec![GradStats::default(); ghist.total_bins()];
for &r in rows {
for f in 0..ghist.n_cols() {
let (fs, fe) = cuts.feature_bins(f);
if let Some(bin) = ghist.feature_bin_at(r as usize, f, fs, fe) {
expected[bin as usize].add(q.row_stats(r as usize));
}
}
}
assert_eq!(q.dequantize(&q.build(ghist, rows)), expected);
}
#[test]
fn histograms_are_exact_for_every_layout_and_thread_count() {
let n = 40_000;
let gpair = gradients(n, 11);
let q = QuantizedGradients::quantize(&gpair, 4, true, 3);
let dense = binned(n, 6, false);
let sparse = binned(n, 6, true);
let all: Vec<u32> = (0..n as u32).collect();
let every_third: Vec<u32> = (0..n as u32).step_by(3).collect();
let small: Vec<u32> = (0..n as u32).step_by(9).take(3000).collect();
let serial = rayon::ThreadPoolBuilder::new()
.num_threads(1)
.build()
.unwrap();
let parallel = rayon::ThreadPoolBuilder::new()
.num_threads(6)
.build()
.unwrap();
for ghist in [&dense, &sparse] {
for rows in [&all, &every_third, &small] {
let a = serial.install(|| q.build(ghist, rows));
let b = parallel.install(|| q.build(ghist, rows));
assert_eq!(a, b);
serial.install(|| assert_matches_reference(ghist, &q, rows));
}
}
let hist = q.build(&dense, &all);
assert_eq!(hist.width(), Width::W64);
let wide = QuantHist::W128(q.build_typed(&dense, &all));
assert_eq!(q.dequantize(&wide), q.dequantize(&hist));
let narrow = q.build(&dense, &small);
assert_eq!(narrow.width(), Width::W32);
let widened = QuantHist::W64(q.build_typed(&dense, &small));
assert_eq!(q.dequantize(&widened), q.dequantize(&narrow));
}
#[test]
fn subtraction_yields_the_sibling_exactly() {
let n = 20_000;
let gpair = gradients(n, 17);
let q = Arc::new(QuantizedGradients::quantize(&gpair, 4, true, 9));
let ghist = binned(n, 5, true);
let all: Vec<u32> = (0..n as u32).collect();
let (left, right): (Vec<u32>, Vec<u32>) = all.iter().partition(|&&r| r % 5 == 0);
let parent = QuantNode {
grads: Arc::clone(&q),
hist: q.build(&ghist, &all),
};
let ((l, lf), (r, rf)) = parent.children(&ghist, &left, &right, Vec::new());
assert_eq!(q.dequantize(&l.hist), q.dequantize(&q.build(&ghist, &left)));
assert_eq!(
q.dequantize(&r.hist),
q.dequantize(&q.build(&ghist, &right))
);
assert_eq!(lf, q.dequantize(&q.build(&ghist, &left)));
assert_eq!(rf, q.dequantize(&q.build(&ghist, &right)));
assert_eq!(l.hist.width(), Width::W32);
assert_eq!(r.hist.width(), Width::W64);
}
}