use super::margins::MarginCaches;
use super::prepare::TrainContext;
use crate::config::{BoosterKind, Dart, TrainingParams};
use crate::model::BoostedModel;
use crate::objective::GradPair;
use crate::rng::Rng;
pub(super) fn round_gradients(
run: &TrainContext,
model: &BoostedModel,
iteration: usize,
margin: &[f32],
gpair: &mut [GradPair],
) -> (Rng, Option<Dropout>) {
let TrainContext {
params,
dtrain,
info,
objective,
..
} = *run;
let mut rng = round_rng(params, iteration, round_salt(params));
let dropout = match params.booster {
BoosterKind::Dart(dart) if dart.has_dropout() => select_dropout(model, &dart, &mut rng),
_ => None,
};
let Some(dropout) = dropout else {
objective.gradient_info_at(margin, info, gpair, iteration);
return (rng, None);
};
let margin_excl = model.predict_margin_dropout(dtrain, &dropout.mask);
objective.gradient_info_at(&margin_excl, info, gpair, iteration);
(rng, Some(dropout))
}
#[derive(Debug, PartialEq)]
pub(super) struct Dropout {
pub(super) mask: Vec<bool>,
pub(super) count: usize,
}
const DART_SALT: u64 = 0x0DA27;
pub(super) fn round_salt(params: &TrainingParams) -> u64 {
match params.booster {
BoosterKind::Dart(dart) if dart.has_dropout() => DART_SALT,
_ => 0,
}
}
pub(super) fn select_dropout(model: &BoostedModel, dart: &Dart, rng: &mut Rng) -> Option<Dropout> {
let existing = model.num_trees();
if existing == 0 {
return None;
}
if dart.skip_drop() > 0.0 && rng.f64() < dart.skip_drop() {
return None;
}
let mut mask = vec![false; existing];
let mut count = 0;
for d in &mut mask {
if rng.f64() < dart.rate_drop() {
*d = true;
count += 1;
}
}
if count == 0 && dart.one_drop() {
mask[rng.range(0..existing)] = true;
count = 1;
}
(count > 0).then_some(Dropout { mask, count })
}
pub(super) fn dart_new_tree_weight(dropout: Option<&Dropout>, params: &TrainingParams) -> f32 {
match dropout {
None => 1.0,
Some(dropout) => 1.0 / (dropout.count as f32 + params.eta as f32),
}
}
pub(super) fn finish_dart(
model: &mut BoostedModel,
params: &TrainingParams,
dropout: &Dropout,
margins: &mut MarginCaches,
) {
let k = dropout.count as f32;
let factor = k / (k + params.eta as f32);
for (i, _) in dropout.mask.iter().enumerate().filter(|&(_, &d)| d) {
model.scale_tree_weight(i, factor);
}
margins.recompute(model);
}
pub(super) fn round_rng(params: &TrainingParams, round: usize, salt: u64) -> Rng {
Rng::new(params.seed ^ (round as u64).wrapping_mul(0x9E37_79B9) ^ salt)
}