use crate::check::{ensure, narrows, positive};
use crate::error::{HessboostError, Result};
use serde::{Deserialize, Serialize};
use std::num::NonZeroUsize;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum AftDistribution {
#[default]
Normal,
Logistic,
Extreme,
}
stored_names! {
AftDistribution { Normal => "normal", Logistic => "logistic", Extreme => "extreme" }
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct PseudoHuber {
slope: f64,
}
impl PseudoHuber {
pub fn new(slope: f64) -> Result<Self> {
positive("huber_slope", slope)?;
let square = (slope as f32) * (slope as f32);
ensure(
"huber_slope",
square.is_finite() && square > 0.0,
format!("squared slope must stay positive and finite in f32, got {slope}"),
)?;
Ok(PseudoHuber { slope })
}
pub fn slope(&self) -> f64 {
self.slope
}
}
impl Default for PseudoHuber {
fn default() -> Self {
PseudoHuber { slope: 1.0 }
}
}
pub(crate) fn validate_alphas(param: &'static str, alphas: &[f64]) -> Result<Vec<f32>> {
let alpha: Vec<f32> = alphas.iter().map(|&a| a as f32).collect();
if alpha.is_empty() {
return Err(HessboostError::invalid_param(
param,
"is required and must list at least one value",
));
}
if !alpha.iter().all(|a| (0.0..=1.0).contains(a)) {
return Err(HessboostError::invalid_param(
param,
"every value must be in the range [0, 1]",
));
}
if !alpha.is_sorted() {
return Err(HessboostError::invalid_param(
param,
"values must be sorted in ascending order",
));
}
Ok(alpha)
}
macro_rules! alpha_list {
($(#[$m:meta])* $ty:ident, $param:literal) => {
$(#[$m])*
#[derive(Debug, Clone, PartialEq)]
pub struct $ty {
alpha: Vec<f64>,
}
impl $ty {
#[doc = concat!("The levels `alpha` (XGBoost `", $param, "`).")]
pub fn new(alpha: impl IntoIterator<Item = f64>) -> Result<Self> {
let alpha: Vec<f64> = alpha.into_iter().collect();
validate_alphas($param, &alpha)?;
Ok($ty { alpha })
}
pub fn alpha(&self) -> &[f64] {
&self.alpha
}
pub(crate) fn alpha_f32(&self) -> Vec<f32> {
self.alpha.iter().map(|&a| a as f32).collect()
}
}
};
}
alpha_list!(
Quantiles,
"quantile_alpha"
);
alpha_list!(
Expectiles,
"expectile_alpha"
);
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Tweedie {
variance_power: f64,
}
impl Tweedie {
pub fn new(variance_power: f64) -> Result<Self> {
let rho = variance_power as f32;
ensure(
"tweedie_variance_power",
variance_power.is_finite() && (1.0f32..2.0).contains(&rho),
format!("must be in [1, 2) (as f32), got {variance_power}"),
)?;
Ok(Tweedie { variance_power })
}
pub fn variance_power(&self) -> f64 {
self.variance_power
}
}
impl Default for Tweedie {
fn default() -> Self {
Tweedie {
variance_power: 1.5,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Aft {
distribution: AftDistribution,
scale: f64,
}
impl Aft {
pub fn new(distribution: AftDistribution, scale: f64) -> Result<Self> {
positive("aft_loss_distribution_scale", scale)?;
narrows("aft_loss_distribution_scale", scale, true)?;
Ok(Aft {
distribution,
scale,
})
}
pub fn with_distribution(distribution: AftDistribution) -> Self {
Aft {
distribution,
scale: 1.0,
}
}
pub fn distribution(&self) -> AftDistribution {
self.distribution
}
pub fn scale(&self) -> f64 {
self.scale
}
}
impl Default for Aft {
fn default() -> Self {
Aft::with_distribution(AftDistribution::Normal)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct RegLoss {
scale_pos_weight: f64,
}
impl RegLoss {
pub fn new(scale_pos_weight: f64) -> Result<Self> {
positive("scale_pos_weight", scale_pos_weight)?;
narrows("scale_pos_weight", scale_pos_weight, true)?;
Ok(RegLoss { scale_pos_weight })
}
pub fn scale_pos_weight(&self) -> f64 {
self.scale_pos_weight
}
}
impl Default for RegLoss {
fn default() -> Self {
RegLoss {
scale_pos_weight: 1.0,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Multiclass {
num_class: usize,
}
impl Multiclass {
pub fn new(num_class: usize) -> Result<Self> {
if num_class < 2 {
return Err(HessboostError::invalid_param(
"num_class",
"multiclass objectives require num_class >= 2",
));
}
Ok(Multiclass { num_class })
}
pub fn num_class(&self) -> usize {
self.num_class
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct LambdaRank {
num_pair_per_sample: NonZeroUsize,
}
impl LambdaRank {
pub fn new(num_pair_per_sample: usize) -> Result<Self> {
NonZeroUsize::new(num_pair_per_sample)
.map(|num_pair_per_sample| LambdaRank {
num_pair_per_sample,
})
.ok_or_else(|| {
HessboostError::invalid_param("lambdarank_num_pair_per_sample", "must be >= 1")
})
}
pub fn num_pair_per_sample(&self) -> usize {
self.num_pair_per_sample.get()
}
}
impl Default for LambdaRank {
fn default() -> Self {
LambdaRank {
num_pair_per_sample: NonZeroUsize::new(32).unwrap_or(NonZeroUsize::MIN),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn refused<T>(built: &Result<T>) -> Option<&'static str> {
match built {
Err(HessboostError::InvalidParameter { name, .. }) => Some(name),
_ => None,
}
}
#[test]
fn constructors_refuse_out_of_range_values_by_xgboost_name() {
for slope in [0.0, -1.0, f64::NAN, 2e19, 1e-30] {
assert_eq!(
refused(&PseudoHuber::new(slope)),
Some("huber_slope"),
"{slope}"
);
}
for rho in [2.0, 2.0 - f64::EPSILON, 0.5, f64::INFINITY] {
assert_eq!(
refused(&Tweedie::new(rho)),
Some("tweedie_variance_power"),
"{rho}"
);
}
assert!(Tweedie::new(1.0).is_ok());
for scale in [0.0, -1.0, f64::INFINITY, f64::NAN, 1e100, 1e-50] {
assert_eq!(
refused(&Aft::new(AftDistribution::Normal, scale)),
Some("aft_loss_distribution_scale"),
"{scale}"
);
}
for alpha in [vec![], vec![0.5, f64::NAN], vec![1.5], vec![0.9, 0.1]] {
assert_eq!(
refused(&Quantiles::new(alpha.clone())),
Some("quantile_alpha"),
"{alpha:?}"
);
assert_eq!(
refused(&Expectiles::new(alpha.clone())),
Some("expectile_alpha"),
"{alpha:?}"
);
}
assert_eq!(
Quantiles::new([0.1, 0.1, 0.9]).unwrap().alpha(),
[0.1, 0.1, 0.9]
);
}
#[test]
fn stored_enum_names_match_serde_and_read_back() {
fn check<T: serde::Serialize + Copy + PartialEq + std::fmt::Debug>(
variants: &[T],
name: fn(T) -> &'static str,
from_name: fn(&str) -> Option<T>,
) {
for &v in variants {
assert_eq!(serde_json::to_value(v).unwrap(), name(v), "{v:?}");
assert_eq!(from_name(name(v)), Some(v));
}
assert_eq!(from_name("no such variant"), None);
}
check(
&[
AftDistribution::Normal,
AftDistribution::Logistic,
AftDistribution::Extreme,
],
AftDistribution::name,
AftDistribution::from_name,
);
check(
&[
crate::objective::distributional::DistGradient::Fisher,
crate::objective::distributional::DistGradient::Hessian,
crate::objective::distributional::DistGradient::Natural,
],
crate::objective::distributional::DistGradient::name,
crate::objective::distributional::DistGradient::from_name,
);
check(
&[
crate::objective::distributional::DistSplitDirection::Random,
crate::objective::distributional::DistSplitDirection::Cyclic,
crate::objective::distributional::DistSplitDirection::All,
],
crate::objective::distributional::DistSplitDirection::name,
crate::objective::distributional::DistSplitDirection::from_name,
);
}
}