use std::ops::Deref;
use std::sync::Arc;
use crate::K_RT_EPS_F32;
use crate::rng::Rng;
#[derive(Debug)]
pub struct ColumnSampler {
tree: Arc<[u32]>,
levels: Vec<Option<Arc<[u32]>>>,
weights: Option<Vec<f32>>,
bylevel: f32,
bynode: f32,
rng: Rng,
seed: u64,
}
#[derive(Debug, Clone)]
pub enum FeatureSet {
Shared(Arc<[u32]>),
Drawn(Vec<u32>),
}
impl Deref for FeatureSet {
type Target = [u32];
#[inline]
fn deref(&self) -> &[u32] {
match self {
FeatureSet::Shared(features) => features,
FeatureSet::Drawn(features) => features,
}
}
}
impl PartialEq for FeatureSet {
fn eq(&self, other: &Self) -> bool {
**self == **other
}
}
impl Eq for FeatureSet {}
impl ColumnSampler {
pub fn new(
n_features: usize,
weights: Option<&[f32]>,
bytree: f64,
bylevel: f64,
bynode: f64,
seed: u64,
) -> Self {
if let Some(w) = weights {
assert_eq!(w.len(), n_features, "one feature weight per feature");
}
let weights = weights.map(<[f32]>::to_vec);
let mut rng = Rng::new(seed);
let all: Vec<u32> = (0..n_features as u32).collect();
let tree = draw(&mut rng, weights.as_deref(), &all, bytree as f32).unwrap_or(all);
ColumnSampler {
tree: tree.into(),
levels: Vec::new(),
weights,
bylevel: bylevel as f32,
bynode: bynode as f32,
rng,
seed,
}
}
pub fn all(n_features: usize) -> Self {
ColumnSampler::new(n_features, None, 1.0, 1.0, 1.0, 0)
}
pub(crate) fn only(features: Vec<u32>, seed: u64) -> Self {
ColumnSampler {
tree: features.into(),
levels: Vec::new(),
weights: None,
bylevel: 1.0,
bynode: 1.0,
rng: Rng::new(seed),
seed,
}
}
pub fn sample(&mut self, depth: usize) -> FeatureSet {
if let Some(features) = self.fixed_features() {
return features;
}
if self.levels.len() <= depth {
self.levels.resize(depth + 1, None);
}
let weights = self.weights.as_deref();
let tree = &self.tree;
let level = self.levels[depth].get_or_insert_with(|| {
draw(&mut self.rng, weights, tree, self.bylevel)
.map_or_else(|| Arc::clone(tree), Arc::from)
});
match draw(&mut self.rng, weights, level, self.bynode) {
Some(drawn) => FeatureSet::Drawn(drawn),
None => FeatureSet::Shared(Arc::clone(level)),
}
}
pub(crate) fn fixed_features(&self) -> Option<FeatureSet> {
(self.bylevel >= 1.0 && self.bynode >= 1.0)
.then(|| FeatureSet::Shared(Arc::clone(&self.tree)))
}
pub(crate) fn seed(&self) -> u64 {
self.seed
}
}
fn draw(rng: &mut Rng, weights: Option<&[f32]>, pool: &[u32], ratio: f32) -> Option<Vec<u32>> {
if ratio >= 1.0 || pool.is_empty() {
return None;
}
let n = ((ratio * pool.len() as f32) as usize).clamp(1, pool.len());
let mut chosen = match weights {
None => {
let mut features = pool.to_vec();
rng.shuffle(&mut features);
features.truncate(n);
features
}
Some(weights) => {
let mut keyed: Vec<(f32, u32)> = pool
.iter()
.map(|&f| {
let w = weights[f as usize].max(K_RT_EPS_F32);
(rng.f32().ln() / w, f)
})
.collect();
keyed.sort_by(|a, b| b.0.total_cmp(&a.0));
keyed.truncate(n);
keyed.into_iter().map(|(_, f)| f).collect()
}
};
chosen.sort_unstable();
Some(chosen)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pass_through_when_ratios_one() {
let mut s = ColumnSampler::all(10);
let all: Vec<u32> = (0..10).collect();
assert_eq!(*s.sample(0), all);
assert_eq!(*s.sample(3), all);
}
#[test]
fn stage_counts_truncate_like_xgboost() {
let mut s = ColumnSampler::new(10, None, 0.75, 0.5, 0.5, 42);
assert_eq!(s.tree.len(), 7);
let f = s.sample(0);
assert_eq!(f.len(), 1);
assert_eq!(s.levels[0].as_ref().unwrap().len(), 3);
let mut tiny = ColumnSampler::new(3, None, 0.01, 0.01, 0.01, 1);
assert_eq!(tiny.sample(0).len(), 1);
}
#[test]
fn level_subset_is_cached_per_depth_and_node_subsets_nest_in_it() {
let mut s = ColumnSampler::new(40, None, 0.8, 0.5, 0.5, 9);
let tree = s.tree.clone();
let first = s.sample(2);
let level = s.levels[2].clone().unwrap();
assert!(level.iter().all(|f| tree.contains(f)));
let mut node_sets = vec![first];
for _ in 0..20 {
node_sets.push(s.sample(2));
}
assert_eq!(s.levels[2].as_ref().unwrap(), &level, "level set is fixed");
for set in &node_sets {
assert!(set.windows(2).all(|w| w[0] < w[1]), "sorted & unique");
assert!(set.iter().all(|f| level.contains(f)));
}
assert!(
node_sets.iter().any(|set| set != &node_sets[0]),
"node subsets are redrawn per call"
);
}
#[test]
fn deterministic_for_seed() {
let weights: Vec<f32> = (0..50).map(|i| (i % 7) as f32).collect();
for w in [None, Some(weights.as_slice())] {
let mut a = ColumnSampler::new(50, w, 0.6, 0.7, 0.8, 7);
let mut b = ColumnSampler::new(50, w, 0.6, 0.7, 0.8, 7);
for depth in [0, 1, 1, 0, 3] {
assert_eq!(a.sample(depth), b.sample(depth));
}
}
}
#[test]
fn single_draw_frequencies_are_proportional_to_weights() {
let weights = [1.0f32, 2.0, 3.0, 4.0, 0.0];
let trials = 40_000;
let mut counts = [0usize; 5];
for seed in 0..trials {
let s = ColumnSampler::new(5, Some(&weights), 0.2, 1.0, 1.0, seed);
counts[s.tree[0] as usize] += 1;
}
assert_eq!(counts[4], 0, "an epsilon weight almost never wins");
for (i, &count) in counts[..4].iter().enumerate() {
let expected = f64::from(weights[i]) / 10.0;
let freq = count as f64 / f64::from(trials as u32);
assert!(
(freq - expected).abs() < 0.012,
"feature {i}: freq {freq} vs {expected}"
);
}
}
#[test]
fn without_replacement_inclusion_matches_successive_sampling() {
let weights = [1.0f32, 1.0, 8.0];
let trials = 30_000u64;
let mut counts = [0usize; 3];
for seed in 0..trials {
let mut s = ColumnSampler::new(3, Some(&weights), 1.0, 1.0, 0.7, seed);
let f = s.sample(0);
assert_eq!(f.len(), 2);
for &x in f.iter() {
counts[x as usize] += 1;
}
}
let freq = |i: usize| counts[i] as f64 / trials as f64;
let heavy = 0.8 + 2.0 * 0.1 * (8.0 / 9.0);
let light = (2.0 - heavy) / 2.0;
assert!((freq(2) - heavy).abs() < 0.01, "heavy {}", freq(2));
assert!((freq(0) - light).abs() < 0.015, "light0 {}", freq(0));
assert!((freq(1) - light).abs() < 0.015, "light1 {}", freq(1));
}
#[test]
fn zero_weight_features_rarely_beat_large_weights() {
let weights = [0.0f32, 5.0, 0.0, 1.0, 0.0];
for seed in 0..200 {
let mut s = ColumnSampler::new(5, Some(&weights), 1.0, 0.4, 1.0, seed);
let two = s.sample(0);
assert_eq!(*two, [1, 3], "large weights win on these seeds");
let mut s = ColumnSampler::new(5, Some(&weights), 0.6, 1.0, 1.0, seed);
let three = s.sample(0);
assert_eq!(three.len(), 3);
assert!(three.contains(&1) && three.contains(&3));
}
}
#[test]
fn zero_and_below_epsilon_weights_are_drawn_alike() {
let trials = 4_000u64;
let mut zero_drawn = 0usize;
for seed in 0..trials {
let s = ColumnSampler::new(2, Some(&[0.0, 1e-8]), 0.5, 1.0, 1.0, seed);
let floored = ColumnSampler::new(2, Some(&[0.0, 0.0]), 0.5, 1.0, 1.0, seed);
assert_eq!(s.tree, floored.tree);
zero_drawn += usize::from(*s.tree == [0]);
}
let freq = zero_drawn as f64 / trials as f64;
assert!((freq - 0.5).abs() < 0.04, "zero-weight frequency {freq}");
}
}