use super::{DistFamily, DistGradient, DistSplitDirection, LOG_LINK_BOUND};
use crate::data::MetaInfo;
use crate::error::{HessboostError, Result};
use crate::objective::{GradPair, Loss, MIN_HESS, SplitGradient, check_label_domain};
use crate::rng::splitmix64;
pub(super) const MIN_CURVATURE: f64 = MIN_HESS as f64;
#[derive(Debug, Clone, Copy)]
pub(crate) struct DistLoss {
family: DistFamily,
gradient: DistGradient,
shared: Option<(DistSplitDirection, u64)>,
}
impl DistLoss {
pub(crate) fn new(family: DistFamily, gradient: DistGradient) -> Self {
DistLoss {
family,
gradient,
shared: None,
}
}
#[must_use]
pub(crate) fn with_split_direction(mut self, direction: DistSplitDirection, seed: u64) -> Self {
self.shared = Some((direction, seed));
self
}
pub(crate) fn split_parameter(&self, iteration: usize) -> Option<usize> {
let k = self.family.n_params();
match self.shared? {
_ if k < 2 => None,
(DistSplitDirection::All, _) => None,
(DistSplitDirection::Cyclic, _) => Some(iteration % k),
(DistSplitDirection::Random, seed) => {
let draw = splitmix64(seed ^ splitmix64(iteration as u64));
Some((draw % k as u64) as usize)
}
}
}
fn row_pairs(self, eta: &[f64], y: f64) -> ([f64; 2], [f64; 2]) {
let g = self.family.gradient(eta, y);
match self.gradient {
DistGradient::Fisher => (g, self.family.fisher(eta).map(|v| v.max(MIN_CURVATURE))),
DistGradient::Hessian => {
let h = self.family.hessian(eta, y);
(g, [h[0][0], h[1][1]].map(|v| v.max(MIN_CURVATURE)))
}
DistGradient::Natural => {
let i = self.family.fisher(eta).map(|v| v.max(MIN_CURVATURE));
([g[0] / i[0], g[1] / i[1]], [1.0; 2])
}
}
}
}
impl Loss for DistLoss {
fn name(&self) -> &str {
self.family.objective_name()
}
fn n_outputs(&self) -> usize {
self.family.n_params()
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let k = self.family.n_params();
let n = labels.len();
crate::objective::rowwise_gradient(
n,
k,
preds,
labels,
weights,
out,
|preds, labels, weights, out| {
for (i, (row, out_row)) in preds
.chunks_exact(k)
.zip(out.chunks_exact_mut(k))
.enumerate()
{
let w = weights.map_or(1.0, |ws| f64::from(ws[i]));
if w == 0.0 {
out_row.fill(GradPair::default());
continue;
}
let mut eta = [0.0; 2];
for (e, &m) in eta.iter_mut().zip(row) {
*e = f64::from(m);
}
let (g, h) = self.row_pairs(&eta[..k], f64::from(labels[i]));
for (j, o) in out_row.iter_mut().enumerate() {
*o = GradPair::new((w * g[j]) as f32, (w * h[j]) as f32);
}
}
},
);
}
fn pred_transform(&self, preds: &mut [f32]) {
let k = self.family.n_params();
for row in preds.chunks_exact_mut(k) {
for (j, v) in row.iter_mut().enumerate() {
if self.family.log_link(j) {
*v = f64::from(*v).clamp(-LOG_LINK_BOUND, LOG_LINK_BOUND).exp() as f32;
}
}
}
}
fn probs_to_margins(&self, scores: &mut [f32]) {
let k = self.family.n_params();
for row in scores.chunks_exact_mut(k) {
for (j, v) in row.iter_mut().enumerate() {
if self.family.log_link(j) {
*v = if *v > 0.0 { v.ln() } else { f32::NAN };
}
}
}
}
fn validate_base_score(&self, base_score: f64) -> Result<()> {
if self.family.n_params() > 1 {
return Err(HessboostError::invalid_param(
"base_score",
"a scalar cannot set the several parameters of a `dist:*` objective; \
supply per-row `base_margin` instead",
));
}
if self.family.log_link(0) {
crate::objective::check_base_score_domain(
base_score,
crate::objective::OutputDomain::Positive,
)?;
}
Ok(())
}
fn base_margins_info(&self, info: &MetaInfo) -> Vec<f32> {
self.family
.mle_margins(info.label_values(), info.weights)
.into_iter()
.map(|m| m as f32)
.collect()
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
check_label_domain(info, |y| self.family.below_support(f64::from(y)))
}
fn default_metric(&self) -> crate::metric::EvalMetric {
crate::metric::EvalMetric::Nll(self.family)
}
fn split_gradient(&self, iteration: usize, gpair: &[GradPair]) -> Option<SplitGradient> {
let m = self.split_parameter(iteration)?;
let k = self.family.n_params();
Some(SplitGradient {
gpair: gpair.iter().skip(m).step_by(k).copied().collect(),
n_targets: 1,
})
}
}