use crate::error::{HessboostError, Result};
use crate::objective::distributional::{DistFamily, DistGradient, DistSplitDirection};
use crate::objective::{AftDistribution, Objective, ObjectiveParts};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq)]
pub struct ModelObjective(Recorded);
#[derive(Debug, Clone, PartialEq)]
enum Recorded {
BuiltIn(Objective),
Other(String),
}
impl ModelObjective {
pub fn new(objective: Objective) -> Result<Self> {
Ok(ModelObjective(match objective {
Objective::Custom(loss) if Objective::is_built_in_name(loss.name()) => {
return Err(HessboostError::invalid_param(
"objective",
format!(
"the custom loss is named `{}`, a built-in objective's name, as \
which a saved model would reload; rename the loss",
loss.name()
),
));
}
Objective::Custom(loss) => Recorded::Other(loss.name().to_owned()),
built_in => Recorded::BuiltIn(built_in),
}))
}
pub fn name(&self) -> &str {
match &self.0 {
Recorded::BuiltIn(objective) => objective.name(),
Recorded::Other(name) => name,
}
}
pub fn built_in(&self) -> Option<&Objective> {
match &self.0 {
Recorded::BuiltIn(objective) => Some(objective),
Recorded::Other(_) => None,
}
}
pub(crate) fn trained_with(objective: &Objective) -> Self {
ModelObjective(match objective {
Objective::Custom(loss) => Recorded::Other(loss.name().to_owned()),
built_in => Recorded::BuiltIn(built_in.clone()),
})
}
pub(crate) fn from_stored(
name: &str,
stored: &StoredObjectiveParams,
num_class: usize,
) -> Result<Self> {
if stored.distribution != DistFamily::from_objective(name) {
return Err(HessboostError::model_format(format!(
"objective parameters name distribution {:?} for objective `{name}`",
stored.distribution
)));
}
let parts = ObjectiveParts {
num_class,
scale_pos_weight: stored.scale_pos_weight,
tweedie_variance_power: stored.tweedie_variance_power,
huber_slope: stored.huber_slope,
lambdarank_num_pair_per_sample: stored.lambdarank_num_pair_per_sample,
quantile_alpha: stored.quantile_alpha.clone(),
expectile_alpha: stored.expectile_alpha.clone(),
aft_loss_distribution: stored.aft_loss_distribution,
aft_loss_distribution_scale: stored.aft_loss_distribution_scale,
dist_gradient: stored.dist_gradient,
dist_split_direction: match stored.dist_split_direction {
DistSplitDirection::Random => None,
direction => Some(direction),
},
};
match Objective::from_parts(name, &parts) {
None => Ok(ModelObjective(Recorded::Other(name.to_owned()))),
Some(Ok(objective)) => Ok(ModelObjective(Recorded::BuiltIn(objective))),
Some(Err(e)) => Err(HessboostError::model_format(format!(
"invalid objective parameters: {e}"
))),
}
}
}
macro_rules! stored_objective_params {
($(#[doc = $sdoc:literal])* ; $(#[doc = $pdoc:literal])* ; $($name:ident: $ty:ty),* $(,)?) => {
$(#[doc = $sdoc])*
#[derive(Debug, Clone, PartialEq, Serialize)]
pub(crate) struct StoredObjectiveParams {
$(pub(crate) $name: $ty,)*
}
$(#[doc = $pdoc])*
#[derive(Deserialize, Default)]
pub(crate) struct PartialStoredObjectiveParams {
$(#[serde(default)] $name: Stored<$ty>,)*
}
impl PartialStoredObjectiveParams {
pub(crate) fn fill(self, objective: &str) -> StoredObjectiveParams {
let d = StoredObjectiveParams::defaults_for(objective);
StoredObjectiveParams {
$($name: self.$name.unwrap_or(d.$name),)*
}
}
}
};
}
stored_objective_params! {
;
;
scale_pos_weight: f64,
max_delta_step: f64,
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,
distribution: Option<DistFamily>,
}
impl StoredObjectiveParams {
pub(crate) fn of(objective: &ModelObjective, max_delta_step: f64) -> Self {
let parts = objective
.built_in()
.map_or_else(ObjectiveParts::default, Objective::parts);
StoredObjectiveParams {
scale_pos_weight: parts.scale_pos_weight,
max_delta_step,
tweedie_variance_power: parts.tweedie_variance_power,
huber_slope: parts.huber_slope,
lambdarank_num_pair_per_sample: parts.lambdarank_num_pair_per_sample,
quantile_alpha: parts.quantile_alpha,
expectile_alpha: parts.expectile_alpha,
aft_loss_distribution: parts.aft_loss_distribution,
aft_loss_distribution_scale: parts.aft_loss_distribution_scale,
dist_gradient: parts.dist_gradient,
dist_split_direction: parts.dist_split_direction.unwrap_or_default(),
distribution: DistFamily::from_objective(objective.name()),
}
}
pub(crate) fn defaults_for(objective: &str) -> Self {
let max_delta_step = if objective == "count:poisson" {
0.7
} else {
0.0
};
let mut defaults = StoredObjectiveParams::of(
&ModelObjective(Recorded::Other(String::new())),
max_delta_step,
);
defaults.distribution = DistFamily::from_objective(objective);
defaults
}
}
#[derive(Default)]
enum Stored<T> {
#[default]
Absent,
Present(T),
}
impl<'de, T: Deserialize<'de>> Deserialize<'de> for Stored<T> {
fn deserialize<D: serde::Deserializer<'de>>(
deserializer: D,
) -> std::result::Result<Self, D::Error> {
T::deserialize(deserializer).map(Stored::Present)
}
}
impl<T> Stored<T> {
fn unwrap_or(self, default: T) -> T {
match self {
Stored::Absent => default,
Stored::Present(value) => value,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::{CustomLoss, GradPair, RegLoss};
#[test]
fn recorded_names_are_never_built_in() {
let loss = |name: &'static str| {
Objective::custom(CustomLoss::new(name, 1, |_, _, _, out: &mut [GradPair]| {
out.fill(GradPair::new(0.0, 1.0));
}))
};
for name in ["reg:squarederror", "reg:linear", "dist:normal"] {
assert!(ModelObjective::new(loss(name)).is_err(), "{name}");
}
let custom = ModelObjective::new(loss("custom:mine")).unwrap();
assert_eq!((custom.name(), custom.built_in()), ("custom:mine", None));
let defaults = StoredObjectiveParams::defaults_for("reg:linear");
let stored = ModelObjective::from_stored("reg:linear", &defaults, 0).unwrap();
assert_eq!(
stored.built_in(),
Some(&Objective::SquaredError(RegLoss::default()))
);
let unknown = ModelObjective::from_stored("rank:foo", &defaults, 0).unwrap();
assert_eq!((unknown.name(), unknown.built_in()), ("rank:foo", None));
}
}