use std::num::NonZeroUsize;
use crate::check::{ensure, fraction, non_negative, positive, unit};
use crate::error::Result;
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct Dart {
rate: f64,
skip: f64,
force_one: bool,
}
impl Dart {
pub fn builder() -> DartBuilder {
DartBuilder {
dart: Dart::default(),
}
}
pub fn rate_drop(&self) -> f64 {
self.rate
}
pub fn skip_drop(&self) -> f64 {
self.skip
}
pub fn one_drop(&self) -> bool {
self.force_one
}
pub(crate) fn has_dropout(&self) -> bool {
self.rate != 0.0 || self.force_one || self.skip != 0.0
}
}
#[derive(Debug, Clone, Copy)]
pub struct DartBuilder {
dart: Dart,
}
impl DartBuilder {
setter!(
rate_drop: f64 => dart.rate
);
setter!(
skip_drop: f64 => dart.skip
);
setter!(
one_drop: bool => dart.force_one
);
pub fn build(self) -> Result<Dart> {
unit("rate_drop", self.dart.rate)?;
unit("skip_drop", self.dart.skip)?;
Ok(self.dart)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct Boulevard {
dropout: f64,
truncation: Option<f64>,
}
impl Boulevard {
pub fn builder() -> BoulevardBuilder {
BoulevardBuilder {
boulevard: Boulevard::default(),
}
}
pub fn dropout(&self) -> f64 {
self.dropout
}
pub fn truncation(&self) -> Option<f64> {
self.truncation
}
}
#[derive(Debug, Clone, Copy)]
pub struct BoulevardBuilder {
boulevard: Boulevard,
}
impl BoulevardBuilder {
setter!(
dropout: f64 => boulevard.dropout
);
setter!(
truncation: f64 => Some(boulevard.truncation)
);
pub fn build(self) -> Result<Boulevard> {
let Boulevard {
dropout,
truncation,
} = self.boulevard;
ensure(
"boulevard_dropout",
dropout.is_finite() && (0.0..1.0).contains(&dropout),
format!("must be in [0, 1), got {dropout}"),
)?;
if let Some(truncation) = truncation {
ensure(
"boulevard_truncation",
truncation.is_finite() && truncation > 0.0,
format!("must be finite and > 0 (leave it unset for none), got {truncation}"),
)?;
}
Ok(self.boulevard)
}
}
const MAX_EBM_OUTER_BAGS: usize = 1024;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct EbmEarlyStopping {
rounds: NonZeroUsize,
tolerance: f64,
}
impl EbmEarlyStopping {
pub const DEFAULT_TOLERANCE: f64 = 1e-5;
pub fn new(rounds: NonZeroUsize, tolerance: f64) -> Result<Self> {
ensure(
"ebm_early_stopping_tolerance",
tolerance.is_finite(),
format!("must be finite, got {tolerance}"),
)?;
Ok(EbmEarlyStopping { rounds, tolerance })
}
pub fn rounds(&self) -> NonZeroUsize {
self.rounds
}
pub fn tolerance(&self) -> f64 {
self.tolerance
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Ebm {
interactions: usize,
outer_bags: usize,
bag_fraction: f64,
boulevard: bool,
early_stopping: Option<EbmEarlyStopping>,
}
impl Default for Ebm {
fn default() -> Self {
Ebm {
interactions: 0,
outer_bags: 1,
bag_fraction: 1.0,
boulevard: false,
early_stopping: None,
}
}
}
impl Ebm {
pub fn builder() -> EbmBuilder {
EbmBuilder {
ebm: Ebm::default(),
}
}
pub fn interactions(&self) -> usize {
self.interactions
}
pub fn outer_bags(&self) -> usize {
self.outer_bags
}
pub fn bag_fraction(&self) -> f64 {
self.bag_fraction
}
pub fn boulevard(&self) -> bool {
self.boulevard
}
pub fn early_stopping(&self) -> Option<EbmEarlyStopping> {
self.early_stopping
}
}
#[derive(Debug, Clone, Copy)]
pub struct EbmBuilder {
ebm: Ebm,
}
impl EbmBuilder {
setter!(
interactions: usize => ebm.interactions
);
setter!(
outer_bags: usize => ebm.outer_bags
);
setter!(
bag_fraction: f64 => ebm.bag_fraction
);
setter!(
boulevard: bool => ebm.boulevard
);
setter!(
early_stopping: EbmEarlyStopping => Some(ebm.early_stopping)
);
pub fn build(self) -> Result<Ebm> {
let e = self.ebm;
ensure(
"ebm_outer_bags",
(1..=MAX_EBM_OUTER_BAGS).contains(&e.outer_bags),
format!("must be in [1, {MAX_EBM_OUTER_BAGS}], got {}", e.outer_bags),
)?;
fraction("ebm_bag_fraction", e.bag_fraction)?;
if e.early_stopping.is_some() {
ensure(
"ebm_early_stopping_rounds",
!e.boulevard,
"a Boulevard EBM averages every round, so it cannot stop at a best round",
)?;
ensure(
"ebm_early_stopping_rounds",
e.bag_fraction < 1.0,
"each bag stops on the rows it does not train on; set `ebm_bag_fraction < 1`",
)?;
}
ensure(
"ebm_outer_bags",
!e.boulevard || (e.outer_bags == 1 && e.bag_fraction == 1.0),
"a bagged Boulevard EBM has no kernel ridge limit its inference covers; use one \
bag of every row",
)?;
Ok(e)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Refresh {
refresh_leaf: bool,
}
impl Refresh {
pub fn stats_only() -> Self {
Refresh {
refresh_leaf: false,
}
}
pub fn refresh_leaf(&self) -> bool {
self.refresh_leaf
}
}
impl Default for Refresh {
fn default() -> Self {
Refresh { refresh_leaf: true }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct QuantizedGrad {
bins: usize,
stochastic_rounding: bool,
renew_leaf: bool,
}
impl QuantizedGrad {
pub fn builder() -> QuantizedGradBuilder {
QuantizedGradBuilder {
quantized: QuantizedGrad::default(),
}
}
pub fn bins(&self) -> usize {
self.bins
}
pub fn stochastic_rounding(&self) -> bool {
self.stochastic_rounding
}
pub fn renew_leaf(&self) -> bool {
self.renew_leaf
}
}
impl Default for QuantizedGrad {
fn default() -> Self {
QuantizedGrad {
bins: 4,
stochastic_rounding: true,
renew_leaf: false,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct QuantizedGradBuilder {
quantized: QuantizedGrad,
}
impl QuantizedGradBuilder {
setter!(
bins: usize => quantized.bins
);
setter!(
stochastic_rounding: bool => quantized.stochastic_rounding
);
setter!(
renew_leaf: bool => quantized.renew_leaf
);
pub fn build(self) -> Result<QuantizedGrad> {
let bins = self.quantized.bins;
ensure(
"num_grad_quant_bins",
(2..=127).contains(&bins),
format!("must be in [2, 127], got {bins}"),
)?;
Ok(self.quantized)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ExtraTrees {
seed: u64,
}
impl ExtraTrees {
pub fn with_seed(seed: u64) -> Self {
ExtraTrees { seed }
}
pub fn seed(&self) -> u64 {
self.seed
}
}
impl Default for ExtraTrees {
fn default() -> Self {
ExtraTrees { seed: 6 }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct LinearTree {
lambda: f64,
}
impl LinearTree {
pub fn new(lambda: f64) -> Result<Self> {
non_negative("linear_lambda", lambda)?;
Ok(LinearTree { lambda })
}
pub fn lambda(&self) -> f64 {
self.lambda
}
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct Langevin {
diffusion_temperature: Option<f64>,
}
impl Langevin {
pub fn builder() -> LangevinBuilder {
LangevinBuilder {
langevin: Langevin::default(),
}
}
pub fn diffusion_temperature(&self) -> Option<f64> {
self.diffusion_temperature
}
}
#[derive(Debug, Clone, Copy)]
pub struct LangevinBuilder {
langevin: Langevin,
}
impl LangevinBuilder {
setter!(
diffusion_temperature: f64 => Some(langevin.diffusion_temperature)
);
pub fn build(self) -> Result<Langevin> {
if let Some(t) = self.langevin.diffusion_temperature {
positive("diffusion_temperature", t)?;
}
Ok(self.langevin)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum ModelShrinkMode {
#[default]
Constant,
Decreasing,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ModelShrink {
rate: f64,
mode: ModelShrinkMode,
}
impl ModelShrink {
pub fn new(rate: f64, mode: ModelShrinkMode) -> Result<Self> {
ensure(
"model_shrink_rate",
rate.is_finite() && rate > 0.0,
format!("must be > 0 (leave model shrinkage unset for none), got {rate}"),
)?;
ensure(
"model_shrink_rate",
mode != ModelShrinkMode::Decreasing || rate < 1.0,
format!("must be in (0, 1) in the decreasing mode, got {rate}"),
)?;
Ok(ModelShrink { rate, mode })
}
pub fn rate(&self) -> f64 {
self.rate
}
pub fn mode(&self) -> ModelShrinkMode {
self.mode
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct QueryBagging {
fraction: f64,
}
impl QueryBagging {
pub fn new(fraction: f64) -> Result<Self> {
ensure(
"bagging_by_query",
fraction > 0.0 && fraction < 1.0,
format!("the fraction of queries kept must be in (0, 1), got {fraction}"),
)?;
Ok(QueryBagging { fraction })
}
pub fn fraction(&self) -> f64 {
self.fraction
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BalancedBagging {
pos: f64,
neg: f64,
}
impl BalancedBagging {
pub fn new(pos: f64, neg: f64) -> Result<Self> {
fraction("pos_bagging_fraction", pos)?;
fraction("neg_bagging_fraction", neg)?;
ensure(
"pos_bagging_fraction",
pos != 1.0 || neg != 1.0,
"with `neg_bagging_fraction` also 1 nothing is bagged; \
leave balanced bagging unset",
)?;
Ok(BalancedBagging { pos, neg })
}
pub fn pos_fraction(&self) -> f64 {
self.pos
}
pub fn neg_fraction(&self) -> f64 {
self.neg
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::HessboostError;
fn refused<T>(built: &Result<T>) -> Option<&'static str> {
match built {
Err(HessboostError::InvalidParameter { name, .. }) => Some(name),
_ => None,
}
}
#[test]
fn groups_refuse_out_of_range_values_by_name() {
for v in [-0.1, 1.5, f64::NAN] {
assert_eq!(
refused(&Dart::builder().rate_drop(v).build()),
Some("rate_drop")
);
assert_eq!(
refused(&Dart::builder().skip_drop(v).build()),
Some("skip_drop")
);
}
for bins in [0, 1, 128] {
assert_eq!(
refused(&QuantizedGrad::builder().bins(bins).build()),
Some("num_grad_quant_bins")
);
}
assert!(QuantizedGrad::builder().bins(127).build().is_ok());
for lambda in [-1.0, f64::INFINITY] {
assert_eq!(refused(&LinearTree::new(lambda)), Some("linear_lambda"));
}
assert!(Refresh::default().refresh_leaf());
assert!(!Refresh::stats_only().refresh_leaf());
assert_eq!(ExtraTrees::default().seed(), 6);
for t in [0.0, -1.0, f64::NAN] {
assert_eq!(
refused(&Langevin::builder().diffusion_temperature(t).build()),
Some("diffusion_temperature")
);
}
for rate in [-0.1, 0.0, f64::NAN, f64::INFINITY] {
assert_eq!(
refused(&ModelShrink::new(rate, ModelShrinkMode::Constant)),
Some("model_shrink_rate")
);
}
assert_eq!(
refused(&ModelShrink::new(1.0, ModelShrinkMode::Decreasing)),
Some("model_shrink_rate")
);
assert!(ModelShrink::new(1.0, ModelShrinkMode::Constant).is_ok());
for v in [0.0, -0.5, 1.1, f64::NAN, f64::INFINITY] {
assert_eq!(
refused(&BalancedBagging::new(v, 0.5)),
Some("pos_bagging_fraction")
);
assert_eq!(
refused(&BalancedBagging::new(0.5, v)),
Some("neg_bagging_fraction")
);
}
assert_eq!(
refused(&BalancedBagging::new(1.0, 1.0)),
Some("pos_bagging_fraction")
);
for fraction in [0.0, 1.0, 1.5, f64::NAN] {
assert_eq!(
refused(&QueryBagging::new(fraction)),
Some("bagging_by_query")
);
}
}
}