use super::super::groups::ModelShrinkMode;
use super::super::params::{
Device, GrowPolicy, Monotone, MultiStrategy, SamplingMethod, TreeMethod,
};
use crate::error::{HessboostError, Result};
use crate::objective::AftDistribution;
use crate::objective::distributional::{DistGradient, DistSplitDirection};
use serde::{Deserialize, Deserializer};
use std::num::NonZeroUsize;
const ALIASES: &[(&str, &str)] = &[
("learning_rate", "eta"),
("min_split_loss", "gamma"),
("reg_lambda", "lambda"),
("reg_alpha", "alpha"),
("random_state", "seed"),
("n_jobs", "nthread"),
];
pub(super) const FIXED: &[(&str, &str)] = &[
("updater", "\"coord_descent\""),
("feature_selector", "\"cyclic\""),
("lambdarank_pair_method", "\"topk\""),
("max_cat_to_onehot", "4"),
("max_cat_threshold", "64"),
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(super) enum FlatBooster {
GbTree,
Dart,
GbLinear,
Boulevard,
Ebm,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(super) enum FlatProcess {
Default,
Update,
}
fn present<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> std::result::Result<Option<T>, D::Error> {
T::deserialize(deserializer).map(Some)
}
#[derive(Debug, Clone, Copy, Deserialize)]
#[serde(from = "usize")]
pub(super) struct FlatLimit(pub(super) Option<NonZeroUsize>);
impl From<usize> for FlatLimit {
fn from(n: usize) -> Self {
FlatLimit(NonZeroUsize::new(n))
}
}
#[derive(Debug, Clone, Copy, Deserialize)]
#[serde(from = "f64")]
pub(super) struct FlatRate(pub(super) Option<f64>);
impl From<f64> for FlatRate {
fn from(rate: f64) -> Self {
FlatRate((rate != 0.0).then_some(rate))
}
}
macro_rules! flat_params {
($($key:ident: $ty:ty,)*) => {
#[derive(Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub(super) struct Flat {
$(
#[serde(default, deserialize_with = "present")]
pub(super) $key: Option<$ty>,
)*
}
pub(super) const KEYS: &[&str] = &[$(stringify!($key)),*];
impl Flat {
pub(super) fn is_set(&self, key: &str) -> bool {
match key {
$(stringify!($key) => self.$key.is_some(),)*
_ => false,
}
}
}
};
}
flat_params! {
booster: FlatBooster,
nthread: FlatLimit,
seed: u64,
device: Device,
objective: String,
num_class: usize,
base_score: Option<f64>,
eval_metric: Vec<String>,
tweedie_variance_power: f64,
huber_slope: f64,
lambdarank_num_pair_per_sample: usize,
quantile_alpha: Vec<f64>,
expectile_alpha: Vec<f64>,
aft_loss_distribution: AftDistribution,
aft_loss_distribution_scale: f64,
dist_gradient: DistGradient,
dist_split_direction: DistSplitDirection,
eta: f64,
gamma: f64,
max_depth: FlatLimit,
max_leaves: FlatLimit,
min_child_weight: f64,
max_delta_step: Option<f64>,
subsample: f64,
colsample_bytree: f64,
colsample_bylevel: f64,
colsample_bynode: f64,
lambda: f64,
alpha: f64,
scale_pos_weight: f64,
tree_method: TreeMethod,
grow_policy: GrowPolicy,
max_bin: usize,
monotone_constraints: Vec<Monotone>,
interaction_constraints: Vec<Vec<u32>>,
num_parallel_tree: usize,
sampling_method: SamplingMethod,
pos_bagging_fraction: f64,
neg_bagging_fraction: f64,
bagging_by_query: bool,
multi_strategy: MultiStrategy,
process_type: FlatProcess,
refresh_leaf: bool,
extra_trees: bool,
extra_seed: u64,
path_smooth: f64,
linear_tree: bool,
linear_lambda: f64,
use_quantized_grad: bool,
num_grad_quant_bins: usize,
stochastic_rounding: bool,
quant_train_renew_leaf: bool,
rate_drop: f64,
skip_drop: f64,
one_drop: bool,
toad_penalty_feature: f64,
toad_penalty_threshold: f64,
langevin: bool,
diffusion_temperature: f64,
model_shrink_rate: FlatRate,
model_shrink_mode: ModelShrinkMode,
posterior_sampling: bool,
boulevard_dropout: f64,
boulevard_truncation: FlatRate,
ebm_interactions: usize,
ebm_outer_bags: usize,
ebm_bag_fraction: f64,
ebm_boulevard: bool,
ebm_early_stopping_rounds: FlatLimit,
ebm_early_stopping_tolerance: f64,
}
fn distance(a: &str, b: &str) -> usize {
let b: Vec<char> = b.chars().collect();
let mut row: Vec<usize> = (0..=b.len()).collect();
for (i, ca) in a.chars().enumerate() {
let mut previous = row[0];
row[0] = i + 1;
for (j, cb) in b.iter().enumerate() {
let substituted = previous + usize::from(ca != *cb);
previous = row[j + 1];
row[j + 1] = substituted.min(row[j] + 1).min(previous + 1);
}
}
row[b.len()]
}
fn unknown_key(key: &str) -> HessboostError {
let candidates = KEYS
.iter()
.copied()
.chain(ALIASES.iter().map(|&(alias, _)| alias))
.chain(FIXED.iter().map(|&(name, _)| name));
let suggestion = candidates
.map(|name| (distance(key, name), name))
.min_by_key(|&(d, name)| (d, name))
.filter(|&(d, _)| d <= (key.len() / 3).max(1))
.map(|(_, name)| name);
HessboostError::Unknown {
kind: "parameter",
name: key.to_owned(),
suggestion,
}
}
pub(super) fn canonical_key(key: &str) -> Result<&'static str> {
if key == "missing" {
return Err(HessboostError::invalid_param(
"missing",
"belongs to the data, not the training parameters: set it when constructing the DMatrix (`DMatrix::from_dense_with_missing`)",
));
}
let key = ALIASES
.iter()
.find(|&&(alias, _)| alias == key)
.map_or(key, |&(_, canonical)| canonical);
KEYS.iter()
.copied()
.find(|&known| known == key)
.ok_or_else(|| unknown_key(key))
}