use std::default::Default;
use super::Interval;
#[derive(Clone, Default)]
pub enum SampleType {
#[default]
Uniform,
Weighted,
}
impl std::fmt::Display for SampleType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let result = match *self {
SampleType::Uniform => "uniform".to_owned(),
SampleType::Weighted => "weighted".to_owned(),
};
write!(f, "{}", result)
}
}
#[derive(Clone, Default)]
pub enum NormalizeType {
#[default]
Tree,
Forest,
}
impl std::fmt::Display for NormalizeType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let result = match *self {
NormalizeType::Tree => "tree".to_owned(),
NormalizeType::Forest => "forest".to_owned(),
};
write!(f, "{}", result)
}
}
#[derive(Builder, Clone)]
#[builder(build_fn(validate = "Self::validate"))]
#[builder(default)]
pub struct DartBoosterParameters {
sample_type: SampleType,
normalize_type: NormalizeType,
rate_drop: f32,
one_drop: bool,
skip_drop: f32,
}
impl Default for DartBoosterParameters {
fn default() -> Self {
DartBoosterParameters {
sample_type: SampleType::default(),
normalize_type: NormalizeType::default(),
rate_drop: 0.0,
one_drop: false,
skip_drop: 0.0,
}
}
}
impl DartBoosterParameters {
pub(crate) fn as_string_pairs(&self) -> Vec<(String, String)> {
vec![
("booster".to_owned(), "dart".to_owned()),
("sample_type".to_owned(), self.sample_type.to_string()),
("normalize_type".to_owned(), self.normalize_type.to_string()),
("rate_drop".to_owned(), self.rate_drop.to_string()),
("one_drop".to_owned(), (self.one_drop as u8).to_string()),
("skip_drop".to_owned(), self.skip_drop.to_string()),
]
}
}
impl DartBoosterParametersBuilder {
fn validate(&self) -> Result<(), String> {
Interval::new_closed_closed(0.0, 1.0).validate(&self.rate_drop, "rate_drop")?;
Interval::new_closed_closed(0.0, 1.0).validate(&self.skip_drop, "skip_drop")?;
Ok(())
}
}