use super::{GradPair, Loss, SplitGradient};
use crate::metric::EvalMetric;
type GradFn = dyn Fn(&[f32], &[f32], Option<&[f32]>, &mut [GradPair]) + Send + Sync;
type TransformFn = dyn Fn(&mut [f32]) + Send + Sync;
type SplitGradFn = dyn Fn(usize, &[GradPair]) -> Option<SplitGradient> + Send + Sync;
pub struct CustomLoss {
name: String,
n_outputs: usize,
base: f32,
default_metric: EvalMetric,
grad_fn: Box<GradFn>,
transform_fn: Option<Box<TransformFn>>,
split_grad_fn: Option<Box<SplitGradFn>>,
}
impl CustomLoss {
pub fn new(
name: impl Into<String>,
n_outputs: usize,
gradient: impl Fn(&[f32], &[f32], Option<&[f32]>, &mut [GradPair]) + Send + Sync + 'static,
) -> Self {
CustomLoss {
name: name.into(),
n_outputs,
base: 0.0,
default_metric: EvalMetric::Rmse,
grad_fn: Box::new(gradient),
transform_fn: None,
split_grad_fn: None,
}
}
#[must_use]
pub fn with_base_margin(mut self, base: f32) -> Self {
self.base = base;
self
}
#[must_use]
pub fn with_default_metric(mut self, metric: EvalMetric) -> Self {
self.default_metric = metric;
self
}
#[must_use]
pub fn with_transform(
mut self,
transform: impl Fn(&mut [f32]) + Send + Sync + 'static,
) -> Self {
self.transform_fn = Some(Box::new(transform));
self
}
#[must_use]
pub fn with_split_gradient(
mut self,
split_grad: impl Fn(usize, &[GradPair]) -> Option<SplitGradient> + Send + Sync + 'static,
) -> Self {
self.split_grad_fn = Some(Box::new(split_grad));
self
}
}
impl Loss for CustomLoss {
fn name(&self) -> &str {
&self.name
}
fn n_outputs(&self) -> usize {
self.n_outputs
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let n_rows = preds.len() / self.n_outputs.max(1);
debug_assert!(labels.len() == n_rows || labels.len() == preds.len());
debug_assert_eq!(out.len(), preds.len());
debug_assert!(weights.is_none_or(|w| w.len() == n_rows));
(self.grad_fn)(preds, labels, weights, out);
}
fn validate_info(&self, info: &crate::data::MetaInfo) -> crate::error::Result<()> {
if info.n_targets() != 1 && info.n_targets() != self.n_outputs {
return Err(crate::error::HessboostError::invalid_data(
"labels",
format!(
"a {}-column label matrix, but custom objective `{}` has {} outputs \
(one label per row or per output expected)",
info.n_targets(),
self.name,
self.n_outputs
),
));
}
Ok(())
}
fn pred_transform(&self, preds: &mut [f32]) {
if let Some(t) = &self.transform_fn {
t(preds);
}
}
fn base_margins_info(&self, _info: &crate::data::MetaInfo) -> Vec<f32> {
vec![self.base; self.n_outputs]
}
fn default_metric(&self) -> EvalMetric {
self.default_metric.clone()
}
fn split_gradient(&self, iteration: usize, gpair: &[GradPair]) -> Option<SplitGradient> {
self.split_grad_fn
.as_ref()
.and_then(|f| f(iteration, gpair))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn custom_squared_error_behaves() {
let obj = CustomLoss::new("custom:sqerr", 1, |preds, labels, w, out| {
for i in 0..preds.len() {
let wi = w.map_or(1.0, |ws| ws[i]);
out[i] = GradPair::new((preds[i] - labels[i]) * wi, wi);
}
});
let mut out = vec![GradPair::default(); 2];
obj.gradient(&[2.0, 0.0], &[1.0, 0.5], None, &mut out);
assert_eq!(out[0], GradPair::new(1.0, 1.0));
assert_eq!(out[1], GradPair::new(-0.5, 1.0));
}
}