use crate::config::{SamplingMethod, TrainingParams};
use crate::data::{DMatrix, GroupInfo};
use crate::rng::Rng;
use crate::tree::builder::all_rows;
use crate::tree::sampler::ColumnSampler;
pub(super) fn sample_rows(
n: usize,
params: &TrainingParams,
meta: RowMeta<'_>,
rng: &mut Rng,
) -> Vec<u32> {
if params.sampling_method == SamplingMethod::GradientBased {
return all_rows(n);
}
let Some(bagging) = params.balanced_bagging else {
return if params.subsample >= 1.0 {
all_rows(n)
} else {
let fraction = params.subsample;
bernoulli_rows(n, |i| i as u32, |_| fraction, n as f64 * fraction, rng)
};
};
let (pos, neg) = (bagging.pos_fraction(), bagging.neg_fraction());
let labels = &meta.labels[..n];
let positives = meta.positives;
let expected = positives as f64 * pos + (n - positives) as f64 * neg;
let fraction = |row: u32| {
if labels[row as usize] == 1.0 {
pos
} else {
neg
}
};
bernoulli_rows(n, |i| i as u32, fraction, expected, rng)
}
fn with_sample_capacity(expected: f64) -> Vec<u32> {
Vec::with_capacity((expected + 4.0 * expected.sqrt() + 16.0) as usize)
}
pub(super) fn bernoulli_rows(
len: usize,
row: impl Fn(usize) -> u32,
keep: impl Fn(u32) -> f64,
expected: f64,
rng: &mut Rng,
) -> Vec<u32> {
let mut kept = with_sample_capacity(expected);
kept.extend((0..len).map(&row).filter(|&r| rng.f64() < keep(r)));
if kept.is_empty() {
kept.push(row(rng.range(0..len)));
}
kept
}
#[derive(Clone, Copy)]
pub(super) struct RowMeta<'a> {
labels: &'a [f32],
positives: usize,
group: Option<&'a GroupInfo>,
}
impl<'a> RowMeta<'a> {
pub(super) fn of(data: &'a DMatrix, params: &TrainingParams) -> Self {
let labels = data.labels().unwrap_or_default();
let positives = if params.balanced_bagging.is_some() {
labels.iter().filter(|&&label| label == 1.0).count()
} else {
0
};
RowMeta {
labels,
positives,
group: data.group(),
}
}
}
pub(super) fn iteration_row_subsets<'a>(
params: &TrainingParams,
per_forest: bool,
meta: RowMeta<'_>,
all: &'a [u32],
rng: &mut Rng,
) -> RoundRows<'a> {
let n = all.len();
let samples = params.sampling_method == SamplingMethod::Uniform
&& (params.subsample < 1.0
|| params.balanced_bagging.is_some()
|| params.bagging_by_query.is_some());
if !samples {
return RoundRows::All(all);
}
let draws = if per_forest {
1
} else {
params.num_parallel_tree
};
let Some(bagging) = params.bagging_by_query else {
return RoundRows::Sampled(
(0..draws)
.map(|_| sample_rows(n, params, meta, rng))
.collect(),
);
};
let queries: Vec<(usize, usize)> = meta
.group
.map_or_else(|| vec![(0, n)], |group| group.iter_ranges().collect());
RoundRows::Sampled(
(0..draws)
.map(|_| {
let fraction = bagging.fraction();
let expected = queries.len() as f64 * fraction;
bernoulli_rows(queries.len(), |q| q as u32, |_| fraction, expected, rng)
.into_iter()
.flat_map(|query| {
let (start, end) = queries[query as usize];
start as u32..end as u32
})
.collect()
})
.collect(),
)
}
#[derive(Debug, PartialEq)]
pub(super) enum RoundRows<'a> {
All(&'a [u32]),
Sampled(Vec<Vec<u32>>),
}
impl RoundRows<'_> {
pub(super) fn rows(&self, parallel: usize) -> &[u32] {
match self {
RoundRows::All(all) => all,
RoundRows::Sampled(subsets) => &subsets[parallel % subsets.len()],
}
}
}
pub(super) fn gradient_sampling(params: &TrainingParams) -> bool {
params.sampling_method == SamplingMethod::GradientBased && params.subsample < 1.0
}
pub(super) fn make_column_sampler(
dtrain: &DMatrix,
params: &TrainingParams,
rng: &mut Rng,
) -> ColumnSampler {
ColumnSampler::new(
dtrain.n_cols(),
dtrain.feature_weights(),
params.colsample_bytree,
params.colsample_bylevel,
params.colsample_bynode,
rng.next_u64(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::{LambdaRank, Objective, RegLoss};
#[test]
fn query_subsets_keep_whole_groups_at_any_thread_count() {
let group_sizes = vec![3, 7, 2, 8, 4, 5, 6];
let group = GroupInfo::from_sizes(&group_sizes);
let n: usize = group_sizes.iter().sum();
let all = all_rows(n);
let params = TrainingParams::builder()
.objective(Objective::RankNdcg(LambdaRank::default()))
.bagging_by_query(crate::config::QueryBagging::new(0.5).unwrap())
.seed(17)
.build()
.unwrap();
let select = |threads| {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap();
pool.install(|| {
iteration_row_subsets(
¶ms,
false,
RowMeta {
labels: &[],
positives: 0,
group: Some(&group),
},
&all,
&mut Rng::new(41),
)
})
};
let one = select(1);
assert_eq!(one, select(4));
let selected = one.rows(0);
assert!(!selected.is_empty() && selected.len() < n);
for (query, (start, end)) in group.iter_ranges().enumerate() {
let included = selected
.iter()
.filter(|&&row| (start..end).contains(&(row as usize)))
.count();
assert!(
included == 0 || included == end - start,
"query {query} split"
);
}
}
#[test]
fn balanced_row_sampler_respects_class_fractions() {
let labels: Vec<f32> = (0..20_000)
.map(|row| if row < 4_000 { 1.0 } else { 0.0 })
.collect();
let params = TrainingParams::builder()
.objective(Objective::BinaryLogistic(RegLoss::default()))
.balanced_bagging(crate::config::BalancedBagging::new(0.6, 0.1).unwrap())
.seed(53)
.build()
.unwrap();
let meta = |labels| RowMeta {
labels,
positives: labels.iter().filter(|&&label| label == 1.0).count(),
group: None,
};
let selected = sample_rows(labels.len(), ¶ms, meta(&labels), &mut Rng::new(77));
let positives = selected
.iter()
.filter(|&&row| labels[row as usize] == 1.0)
.count();
let negatives = selected.len() - positives;
let pos_rate = positives as f64 / 4_000.0;
let neg_rate = negatives as f64 / 16_000.0;
assert!((0.57..0.63).contains(&pos_rate), "positive rate {pos_rate}");
assert!(
(0.092..0.108).contains(&neg_rate),
"negative rate {neg_rate}"
);
let negatives_only = vec![0.0f32; 20_000];
let kept = sample_rows(20_000, ¶ms, meta(&negatives_only), &mut Rng::new(77)).len();
assert!(
(1_840..2_160).contains(&kept),
"kept {kept} of 20000 negatives"
);
}
}