use super::{GradPair, Objective};
type GradFn = dyn Fn(&[f32], &[f32], Option<&[f32]>, &mut [GradPair]) + Send + Sync;
type TransformFn = dyn Fn(&mut [f32]) + Send + Sync;
pub struct CustomObjective {
name: String,
n_outputs: usize,
base: f32,
default_metric: String,
grad_fn: Box<GradFn>,
transform_fn: Option<Box<TransformFn>>,
}
impl CustomObjective {
pub fn new(
name: impl Into<String>,
n_outputs: usize,
base: f32,
default_metric: impl Into<String>,
grad_fn: impl Fn(&[f32], &[f32], Option<&[f32]>, &mut [GradPair]) + Send + Sync + 'static,
) -> Self {
CustomObjective {
name: name.into(),
n_outputs,
base,
default_metric: default_metric.into(),
grad_fn: Box::new(grad_fn),
transform_fn: None,
}
}
#[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
}
}
impl Objective for CustomObjective {
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],
) {
super::check_gradient_inputs(labels.len(), self.n_outputs, preds, labels, weights, out);
(self.grad_fn)(preds, labels, weights, out);
}
fn pred_transform(&self, preds: &mut [f32]) {
if let Some(t) = &self.transform_fn {
t(preds);
}
}
fn base_margins(
&self,
_labels: &[f32],
_weights: Option<&[f32]>,
_group: Option<&crate::data::GroupInfo>,
) -> Vec<f32> {
vec![self.base; self.n_outputs]
}
fn default_metric(&self) -> String {
self.default_metric.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn custom_squared_error_behaves() {
let obj = CustomObjective::new("custom:sqerr", 1, 0.0, "rmse", |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));
}
}