use super::{GradPair, Loss};
use crate::data::{Labels, MetaInfo};
use crate::error::Result;
use std::num::NonZeroUsize;
pub(crate) struct MultiTarget {
inner: Box<dyn Loss>,
n_targets: usize,
}
impl MultiTarget {
pub(crate) fn new(inner: Box<dyn Loss>, n_targets: usize) -> Self {
debug_assert_eq!(inner.n_outputs(), 1);
MultiTarget { inner, n_targets }
}
fn cells<'a>(info: &MetaInfo<'a>, cell_weights: Option<&'a [f32]>) -> MetaInfo<'a> {
MetaInfo::new(info.label_values(), cell_weights, None)
}
}
impl Loss for MultiTarget {
fn name(&self) -> &str {
self.inner.name()
}
fn n_outputs(&self) -> usize {
self.n_targets
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let n_targets = NonZeroUsize::new(self.n_targets).expect("a label matrix has targets");
let info = MetaInfo {
n_rows: labels.len() / self.n_targets,
labels: Some(Labels::new(labels, n_targets)),
..MetaInfo::new(labels, weights, None)
};
self.gradient_info(preds, &info, out);
}
fn gradient_info(&self, preds: &[f32], info: &MetaInfo, out: &mut [GradPair]) {
let Ok(cell_weights) = info.cell_weights() else {
out.fill(GradPair::default());
return;
};
self.inner
.gradient_info(preds, &Self::cells(info, cell_weights.as_deref()), out);
}
fn const_hess(&self) -> bool {
self.inner.const_hess()
}
fn pred_transform(&self, preds: &mut [f32]) {
self.inner.pred_transform(preds);
}
fn base_margins_info(&self, info: &MetaInfo) -> Vec<f32> {
let mut column = Vec::with_capacity(info.n_rows);
(0..self.n_targets)
.map(|j| {
column.clear();
column.extend(info.label_values().iter().skip(j).step_by(self.n_targets));
self.inner
.base_margins_info(&MetaInfo::new(&column, info.weights, None))[0]
})
.collect()
}
fn eval_transform(&self, preds: &mut [f32]) {
self.inner.eval_transform(preds);
}
fn probs_to_margins(&self, scores: &mut [f32]) {
self.inner.probs_to_margins(scores);
}
fn validate_base_score(&self, base_score: f64) -> Result<()> {
self.inner.validate_base_score(base_score)
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
super::check_label_width(info, self.n_targets)?;
let cell_weights = info.cell_weights()?;
self.inner
.validate_info(&Self::cells(info, cell_weights.as_deref()))
}
fn requires_labels(&self) -> bool {
self.inner.requires_labels()
}
fn default_metric(&self) -> crate::metric::EvalMetric {
self.inner.default_metric()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::TrainingParams;
use crate::error::HessboostError;
use crate::objective::base_margins;
use serde_json::json;
use std::sync::Arc;
fn objective(name: &str, n_targets: usize) -> Arc<dyn Loss> {
let mut flat = vec![("objective", json!(name))];
if matches!(name, "reg:squarederror" | "binary:logistic") {
flat.push(("scale_pos_weight", json!(2.0)));
}
let params = TrainingParams::from_xgboost(flat).unwrap();
params.loss(n_targets).unwrap()
}
fn columns() -> (Vec<f32>, [Vec<f32>; 2], Vec<f32>) {
let a: Vec<f32> = (0..7).map(|i| (i % 2) as f32).collect();
let b: Vec<f32> = (0..7).map(|i| f32::from(u8::from(i % 3 == 0))).collect();
let matrix = a.iter().zip(&b).flat_map(|(&x, &y)| [x, y]).collect();
let weights = (0..7).map(|i| 0.5 + i as f32 * 0.25).collect();
(matrix, [a, b], weights)
}
#[test]
fn each_output_is_the_single_target_objective_on_its_column() {
let (matrix, cols, weights) = columns();
for name in [
"reg:squarederror",
"reg:pseudohubererror",
"reg:logistic",
"binary:logistic",
] {
let multi = objective(name, 2);
let single = objective(name, 1);
assert_eq!(multi.n_outputs(), 2);
assert_eq!(multi.name(), single.name());
let preds: Vec<f32> = (0..14).map(|i| i as f32 * 0.3 - 2.0).collect();
for w in [None, Some(weights.as_slice())] {
let info = MetaInfo {
n_rows: 7,
labels: Some(crate::data::Labels::new(
&matrix,
std::num::NonZeroUsize::new(2).unwrap(),
)),
..MetaInfo::new(&matrix, w, None)
};
let mut out = vec![GradPair::default(); 14];
multi.gradient_info(&preds, &info, &mut out);
let margins = multi.base_margins_info(&info);
for (j, col) in cols.iter().enumerate() {
let col_preds: Vec<f32> = preds.iter().skip(j).step_by(2).copied().collect();
let mut expected = vec![GradPair::default(); 7];
single.gradient(&col_preds, col, w, &mut expected);
for (row, e) in expected.iter().enumerate() {
let got = out[row * 2 + j];
assert_eq!(got.grad.to_bits(), e.grad.to_bits(), "{name} ({row},{j})");
assert_eq!(got.hess.to_bits(), e.hess.to_bits(), "{name} ({row},{j})");
}
let intercept = base_margins(single.as_ref(), col, w);
assert_eq!(margins[j].to_bits(), intercept[0].to_bits(), "{name} {j}");
}
}
}
}
#[test]
fn label_domain_checks_every_cell() {
let multi = objective("binary:logistic", 2);
let labels = [0.0, 1.0, 1.0, 1.5];
let info = MetaInfo {
n_rows: 2,
labels: Some(crate::data::Labels::new(
&labels,
std::num::NonZeroUsize::new(2).unwrap(),
)),
..MetaInfo::new(&labels, None, None)
};
assert!(multi.validate_info(&info).is_err());
let info = MetaInfo {
labels: Some(crate::data::Labels::new(
&[0.0, 1.0, 1.0, 0.5],
std::num::NonZeroUsize::new(2).unwrap(),
)),
..info
};
assert!(multi.validate_info(&info).is_ok());
}
#[test]
fn rejects_a_different_label_width() {
let multi = objective("reg:squarederror", 2);
let labels = [0.0, 1.0];
let single = MetaInfo::new(&labels, None, None);
assert!(matches!(
multi.validate_info(&single),
Err(HessboostError::InvalidData {
input: "labels",
..
})
));
}
#[test]
fn inconsistent_metadata_is_refused_without_allocating() {
let multi = objective("reg:squarederror", 2);
let labels = [0.0, 1.0, 1.0, 0.0];
let info = MetaInfo {
n_rows: 2,
labels: Some(crate::data::Labels::new(
&labels,
std::num::NonZeroUsize::new(2).unwrap(),
)),
..MetaInfo::new(&labels, Some(&[1.0]), None)
};
assert!(matches!(
multi.validate_info(&info),
Err(HessboostError::InvalidData {
input: "weights",
..
})
));
let huge = MetaInfo {
labels: Some(crate::data::Labels::new(
&labels,
std::num::NonZeroUsize::MAX,
)),
weights: Some(&[1.0, 1.0]),
..info
};
let mut out = vec![GradPair::new(1.0, 1.0); 4];
multi.gradient_info(&[0.5; 4], &huge, &mut out);
assert_eq!(out, vec![GradPair::default(); 4]);
}
}