use super::distributional::{
DistFamily, DistGradient, DistLoss, DistSplitDirection, Distributional,
};
use super::{
AbsoluteError, Aft, AftDistribution, AftLoss, Cox, Expectile, Expectiles, Gamma, Hinge,
LambdaMart, LambdaRank, LogisticLoss, Loss, Multiclass, Poisson, PseudoHuber, PseudoHuberLoss,
Quantile, Quantiles, RegLoss, Softmax, SquaredError, SquaredLogError, Tweedie, TweedieLoss,
Xendcg, multi_target::MultiTarget,
};
use crate::error::{HessboostError, Result};
use serde::Serialize;
use serde_json::Value;
use std::fmt;
use std::sync::Arc;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum Objective {
SquaredError(RegLoss),
SquaredLogError,
PseudoHuber(PseudoHuber),
AbsoluteError,
Quantile(Quantiles),
Expectile(Expectiles),
RegLogistic(RegLoss),
BinaryLogistic(RegLoss),
BinaryLogitRaw(RegLoss),
BinaryHinge,
Softmax(Multiclass),
Softprob(Multiclass),
Poisson,
Gamma(RegLoss),
Tweedie(Tweedie),
RankPairwise(LambdaRank),
RankNdcg(LambdaRank),
RankMap(LambdaRank),
RankXendcg,
Cox,
Aft(Aft),
Dist(Distributional),
Custom(Arc<dyn Loss>),
}
impl Default for Objective {
fn default() -> Self {
Objective::SquaredError(RegLoss::default())
}
}
impl fmt::Debug for dyn Loss + '_ {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Loss").field("name", &self.name()).finish()
}
}
impl PartialEq for Objective {
fn eq(&self, other: &Self) -> bool {
use Objective as O;
match (self, other) {
(O::SquaredLogError, O::SquaredLogError)
| (O::AbsoluteError, O::AbsoluteError)
| (O::BinaryHinge, O::BinaryHinge)
| (O::Poisson, O::Poisson)
| (O::RankXendcg, O::RankXendcg)
| (O::Cox, O::Cox) => true,
(O::PseudoHuber(a), O::PseudoHuber(b)) => a == b,
(O::Quantile(a), O::Quantile(b)) => a == b,
(O::Expectile(a), O::Expectile(b)) => a == b,
(O::SquaredError(a), O::SquaredError(b))
| (O::RegLogistic(a), O::RegLogistic(b))
| (O::BinaryLogistic(a), O::BinaryLogistic(b))
| (O::BinaryLogitRaw(a), O::BinaryLogitRaw(b))
| (O::Gamma(a), O::Gamma(b)) => a == b,
(O::Softmax(a), O::Softmax(b)) | (O::Softprob(a), O::Softprob(b)) => a == b,
(O::Tweedie(a), O::Tweedie(b)) => a == b,
(O::RankPairwise(a), O::RankPairwise(b))
| (O::RankNdcg(a), O::RankNdcg(b))
| (O::RankMap(a), O::RankMap(b)) => a == b,
(O::Aft(a), O::Aft(b)) => a == b,
(O::Dist(a), O::Dist(b)) => a == b,
(O::Custom(a), O::Custom(b)) => Arc::ptr_eq(a, b),
_ => false,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum LabelMatrix {
PerColumn,
Native,
Refused,
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct LossContext {
pub(crate) n_targets: usize,
pub(crate) max_delta_step: f64,
pub(crate) shared_tree_seed: Option<u64>,
pub(crate) seed: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct ObjectiveParts {
pub(crate) num_class: usize,
pub(crate) scale_pos_weight: f64,
pub(crate) tweedie_variance_power: f64,
pub(crate) huber_slope: f64,
pub(crate) lambdarank_num_pair_per_sample: usize,
pub(crate) quantile_alpha: Vec<f64>,
pub(crate) expectile_alpha: Vec<f64>,
pub(crate) aft_loss_distribution: AftDistribution,
pub(crate) aft_loss_distribution_scale: f64,
pub(crate) dist_gradient: DistGradient,
pub(crate) dist_split_direction: Option<DistSplitDirection>,
}
impl Default for ObjectiveParts {
fn default() -> Self {
ObjectiveParts {
num_class: 0,
scale_pos_weight: RegLoss::default().scale_pos_weight(),
tweedie_variance_power: Tweedie::default().variance_power(),
huber_slope: PseudoHuber::default().slope(),
lambdarank_num_pair_per_sample: LambdaRank::default().num_pair_per_sample(),
quantile_alpha: Vec::new(),
expectile_alpha: Vec::new(),
aft_loss_distribution: AftDistribution::Normal,
aft_loss_distribution_scale: Aft::default().scale(),
dist_gradient: DistGradient::Fisher,
dist_split_direction: None,
}
}
}
pub(crate) struct ObjectiveParam {
pub(crate) key: &'static str,
pub(crate) users: &'static str,
pub(crate) read_by: fn(&Objective) -> bool,
pub(crate) value: fn(&ObjectiveParts) -> Option<Value>,
}
fn json(value: impl Serialize) -> Option<Value> {
serde_json::to_value(value).ok()
}
pub(crate) const OBJECTIVE_PARAMS: &[ObjectiveParam] = &[
ObjectiveParam {
key: "num_class",
users: "`multi:softmax` and `multi:softprob`",
read_by: |o| matches!(o, Objective::Softmax(_) | Objective::Softprob(_)),
value: |p| json(p.num_class),
},
ObjectiveParam {
key: "scale_pos_weight",
users: "`reg:squarederror`, `reg:gamma`, `reg:logistic`, `binary:logistic`, and \
`binary:logitraw`",
read_by: |o| {
matches!(
o,
Objective::SquaredError(_)
| Objective::RegLogistic(_)
| Objective::BinaryLogistic(_)
| Objective::BinaryLogitRaw(_)
| Objective::Gamma(_)
)
},
value: |p| json(p.scale_pos_weight),
},
ObjectiveParam {
key: "tweedie_variance_power",
users: "`reg:tweedie`",
read_by: |o| matches!(o, Objective::Tweedie(_)),
value: |p| json(p.tweedie_variance_power),
},
ObjectiveParam {
key: "huber_slope",
users: "`reg:pseudohubererror` and the `mphe` metric",
read_by: |o| matches!(o, Objective::PseudoHuber(_)),
value: |p| json(p.huber_slope),
},
ObjectiveParam {
key: "lambdarank_num_pair_per_sample",
users: "the `rank:*` objectives",
read_by: |o| {
matches!(
o,
Objective::RankPairwise(_) | Objective::RankNdcg(_) | Objective::RankMap(_)
)
},
value: |p| json(p.lambdarank_num_pair_per_sample),
},
ObjectiveParam {
key: "quantile_alpha",
users: "`reg:quantileerror` and the `quantile` metric",
read_by: |o| matches!(o, Objective::Quantile(_)),
value: |p| json(&p.quantile_alpha),
},
ObjectiveParam {
key: "expectile_alpha",
users: "`reg:expectileerror` and the `expectile` metric",
read_by: |o| matches!(o, Objective::Expectile(_)),
value: |p| json(&p.expectile_alpha),
},
ObjectiveParam {
key: "aft_loss_distribution",
users: "`survival:aft` and the `aft-nloglik` metric",
read_by: |o| matches!(o, Objective::Aft(_)),
value: |p| json(p.aft_loss_distribution),
},
ObjectiveParam {
key: "aft_loss_distribution_scale",
users: "`survival:aft` and the `aft-nloglik` metric",
read_by: |o| matches!(o, Objective::Aft(_)),
value: |p| json(p.aft_loss_distribution_scale),
},
ObjectiveParam {
key: "dist_gradient",
users: "the `dist:*` objectives",
read_by: |o| matches!(o, Objective::Dist(_)),
value: |p| json(p.dist_gradient),
},
ObjectiveParam {
key: "dist_split_direction",
users: "the `dist:*` objectives",
read_by: |o| matches!(o, Objective::Dist(_)),
value: |p| json(p.dist_split_direction?),
},
];
impl Objective {
pub fn custom(loss: impl Loss + 'static) -> Self {
Objective::Custom(Arc::new(loss))
}
pub fn name(&self) -> &str {
match self {
Objective::SquaredError(_) => "reg:squarederror",
Objective::SquaredLogError => "reg:squaredlogerror",
Objective::PseudoHuber(_) => "reg:pseudohubererror",
Objective::AbsoluteError => "reg:absoluteerror",
Objective::Quantile(_) => "reg:quantileerror",
Objective::Expectile(_) => "reg:expectileerror",
Objective::RegLogistic(_) => "reg:logistic",
Objective::BinaryLogistic(_) => "binary:logistic",
Objective::BinaryLogitRaw(_) => "binary:logitraw",
Objective::BinaryHinge => "binary:hinge",
Objective::Softmax(_) => "multi:softmax",
Objective::Softprob(_) => "multi:softprob",
Objective::Poisson => "count:poisson",
Objective::Gamma(_) => "reg:gamma",
Objective::Tweedie(_) => "reg:tweedie",
Objective::RankPairwise(_) => "rank:pairwise",
Objective::RankNdcg(_) => "rank:ndcg",
Objective::RankMap(_) => "rank:map",
Objective::RankXendcg => "rank:xendcg",
Objective::Cox => "survival:cox",
Objective::Aft(_) => "survival:aft",
Objective::Dist(dist) => dist.family().objective_name(),
Objective::Custom(loss) => loss.name(),
}
}
pub fn num_class(&self) -> Option<usize> {
match self {
Objective::Softmax(classes) | Objective::Softprob(classes) => Some(classes.num_class()),
_ => None,
}
}
pub fn dist_family(&self) -> Option<DistFamily> {
match self {
Objective::Dist(dist) => Some(dist.family()),
_ => None,
}
}
pub(crate) fn default_max_delta_step(&self) -> f64 {
if matches!(self, Objective::Poisson) {
0.7
} else {
0.0
}
}
pub(crate) fn has_adaptive_leaves(&self) -> bool {
matches!(self, Objective::AbsoluteError | Objective::Quantile(_))
}
pub(crate) fn is_unweighted_squared_error(&self) -> bool {
matches!(self, Objective::SquaredError(r) if *r == RegLoss::default())
}
pub(crate) fn is_ranking(&self) -> bool {
matches!(
self,
Objective::RankPairwise(_)
| Objective::RankNdcg(_)
| Objective::RankMap(_)
| Objective::RankXendcg
)
}
pub(crate) fn is_binary_classifier(&self) -> bool {
matches!(
self,
Objective::BinaryLogistic(_) | Objective::BinaryLogitRaw(_) | Objective::BinaryHinge
)
}
pub(crate) fn predicts_class_index(&self) -> bool {
matches!(self, Objective::Softmax(_))
}
pub(crate) fn label_matrix(&self) -> LabelMatrix {
match self {
Objective::SquaredError(_)
| Objective::PseudoHuber(_)
| Objective::RegLogistic(_)
| Objective::BinaryLogistic(_) => LabelMatrix::PerColumn,
Objective::AbsoluteError | Objective::Custom(_) => LabelMatrix::Native,
Objective::SquaredLogError
| Objective::Quantile(_)
| Objective::Expectile(_)
| Objective::BinaryLogitRaw(_)
| Objective::BinaryHinge
| Objective::Softmax(_)
| Objective::Softprob(_)
| Objective::Poisson
| Objective::Gamma(_)
| Objective::Tweedie(_)
| Objective::RankPairwise(_)
| Objective::RankNdcg(_)
| Objective::RankMap(_)
| Objective::RankXendcg
| Objective::Cox
| Objective::Aft(_)
| Objective::Dist(_) => LabelMatrix::Refused,
}
}
pub(crate) fn is_built_in_name(name: &str) -> bool {
Objective::from_parts(name, &ObjectiveParts::default()).is_some()
}
pub(crate) fn from_parts(name: &str, parts: &ObjectiveParts) -> Option<Result<Objective>> {
let reg_loss = || RegLoss::new(parts.scale_pos_weight);
let classes = || Multiclass::new(parts.num_class);
let rank = || LambdaRank::new(parts.lambdarank_num_pair_per_sample);
let objective = match name {
"reg:squarederror" | "reg:linear" => reg_loss().map(Objective::SquaredError),
"reg:squaredlogerror" => Ok(Objective::SquaredLogError),
"reg:pseudohubererror" => {
PseudoHuber::new(parts.huber_slope).map(Objective::PseudoHuber)
}
"reg:absoluteerror" => Ok(Objective::AbsoluteError),
"reg:quantileerror" => {
Quantiles::new(parts.quantile_alpha.iter().copied()).map(Objective::Quantile)
}
"reg:expectileerror" => {
Expectiles::new(parts.expectile_alpha.iter().copied()).map(Objective::Expectile)
}
"reg:logistic" => reg_loss().map(Objective::RegLogistic),
"binary:logistic" => reg_loss().map(Objective::BinaryLogistic),
"binary:logitraw" => reg_loss().map(Objective::BinaryLogitRaw),
"binary:hinge" => Ok(Objective::BinaryHinge),
"multi:softmax" => classes().map(Objective::Softmax),
"multi:softprob" => classes().map(Objective::Softprob),
"count:poisson" => Ok(Objective::Poisson),
"reg:gamma" => reg_loss().map(Objective::Gamma),
"reg:tweedie" => Tweedie::new(parts.tweedie_variance_power).map(Objective::Tweedie),
"rank:pairwise" => rank().map(Objective::RankPairwise),
"rank:ndcg" => rank().map(Objective::RankNdcg),
"rank:map" => rank().map(Objective::RankMap),
"rank:xendcg" => Ok(Objective::RankXendcg),
"survival:cox" => Ok(Objective::Cox),
"survival:aft" => Aft::new(
parts.aft_loss_distribution,
parts.aft_loss_distribution_scale,
)
.map(Objective::Aft),
other => {
let family = DistFamily::from_objective(other)?;
let dist = Distributional::new(family).with_gradient(parts.dist_gradient);
Ok(Objective::Dist(match parts.dist_split_direction {
Some(direction) => dist.with_split_direction(direction),
None => dist,
}))
}
};
Some(objective)
}
pub(crate) fn parts(&self) -> ObjectiveParts {
let d = ObjectiveParts::default();
match self {
Objective::PseudoHuber(huber) => ObjectiveParts {
huber_slope: huber.slope(),
..d
},
Objective::Quantile(q) => ObjectiveParts {
quantile_alpha: q.alpha().to_vec(),
..d
},
Objective::Expectile(e) => ObjectiveParts {
expectile_alpha: e.alpha().to_vec(),
..d
},
Objective::SquaredError(r)
| Objective::RegLogistic(r)
| Objective::BinaryLogistic(r)
| Objective::BinaryLogitRaw(r)
| Objective::Gamma(r) => ObjectiveParts {
scale_pos_weight: r.scale_pos_weight(),
..d
},
Objective::Softmax(c) | Objective::Softprob(c) => ObjectiveParts {
num_class: c.num_class(),
..d
},
Objective::Tweedie(t) => ObjectiveParts {
tweedie_variance_power: t.variance_power(),
..d
},
Objective::RankPairwise(r) | Objective::RankNdcg(r) | Objective::RankMap(r) => {
ObjectiveParts {
lambdarank_num_pair_per_sample: r.num_pair_per_sample(),
..d
}
}
Objective::Aft(aft) => ObjectiveParts {
aft_loss_distribution: aft.distribution(),
aft_loss_distribution_scale: aft.scale(),
..d
},
Objective::Dist(dist) => ObjectiveParts {
dist_gradient: dist.gradient(),
dist_split_direction: dist.split_direction(),
..d
},
Objective::SquaredLogError
| Objective::AbsoluteError
| Objective::BinaryHinge
| Objective::Poisson
| Objective::RankXendcg
| Objective::Cox
| Objective::Custom(_) => d,
}
}
pub(crate) fn build_loss(&self, context: &LossContext) -> Result<Arc<dyn Loss>> {
let n_targets = context.n_targets;
let single: Box<dyn Loss> = match self {
Objective::Custom(loss) => return Ok(Arc::clone(loss)),
Objective::AbsoluteError => return Ok(Arc::new(AbsoluteError::new(n_targets))),
Objective::SquaredError(r) => Box::new(SquaredError::new(r.scale_pos_weight() as f32)),
Objective::SquaredLogError => Box::new(SquaredLogError),
Objective::PseudoHuber(huber) => Box::new(PseudoHuberLoss::new(*huber)),
Objective::Quantile(q) => Box::new(Quantile::from_levels(q.clone())),
Objective::Expectile(e) => Box::new(Expectile::from_levels(e.clone())),
Objective::RegLogistic(r) => {
Box::new(LogisticLoss::regression(r.scale_pos_weight() as f32))
}
Objective::BinaryLogistic(r) => {
Box::new(LogisticLoss::new(r.scale_pos_weight() as f32))
}
Objective::BinaryLogitRaw(r) => {
Box::new(LogisticLoss::raw(r.scale_pos_weight() as f32))
}
Objective::BinaryHinge => Box::new(Hinge),
Objective::Softmax(c) => Box::new(Softmax::new(c.num_class(), false)),
Objective::Softprob(c) => Box::new(Softmax::new(c.num_class(), true)),
Objective::Poisson => Box::new(Poisson::new(context.max_delta_step as f32)),
Objective::Gamma(r) => Box::new(Gamma::new(r.scale_pos_weight() as f32)),
Objective::Tweedie(t) => Box::new(TweedieLoss::new(*t)),
Objective::RankPairwise(r) => Box::new(LambdaMart::pairwise(r.num_pair_per_sample())),
Objective::RankNdcg(r) => Box::new(LambdaMart::ndcg(r.num_pair_per_sample())),
Objective::RankMap(r) => Box::new(LambdaMart::map(r.num_pair_per_sample())),
Objective::RankXendcg => Box::new(Xendcg::new(context.seed)),
Objective::Cox => Box::new(Cox),
Objective::Aft(aft) => Box::new(AftLoss::new(aft.distribution(), aft.scale() as f32)),
Objective::Dist(dist) => {
let loss = DistLoss::new(dist.family(), dist.gradient());
Box::new(match context.shared_tree_seed {
Some(seed) => {
loss.with_split_direction(dist.split_direction().unwrap_or_default(), seed)
}
None => loss,
})
}
};
if n_targets <= 1 {
return Ok(Arc::from(single));
}
match self.label_matrix() {
LabelMatrix::PerColumn => Ok(Arc::new(MultiTarget::new(single, n_targets))),
LabelMatrix::Native | LabelMatrix::Refused => Err(HessboostError::invalid_data(
"labels",
format!(
"objective `{}` supports one target per row, got {n_targets}",
self.name()
),
)),
}
}
}