use super::{Metric, consistent};
use crate::data::MetaInfo;
use crate::error::{HessboostError, Result};
use crate::objective::AftDistribution;
use crate::objective::{abs_label_order, aft_nloglik};
use rayon::prelude::*;
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) struct CoxNLogLik;
impl Metric for CoxNLogLik {
fn name(&self) -> &'static str {
"cox-nloglik"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
nan_unless_consistent!(preds, labels, weights, 1);
let n = labels.len();
let order = abs_label_order(labels);
let mut exp_p_sum: f64 = preds[..n].iter().map(|&p| f64::from(p)).sum();
let mut out = 0.0f64;
let mut accumulated_sum = 0.0f64;
let mut num_events = 0u64;
for (i, &ind) in order.iter().enumerate() {
let label = labels[ind];
if label > 0.0 {
out -= f64::from(preds[ind].ln()) - exp_p_sum.ln();
num_events += 1;
}
accumulated_sum += f64::from(preds[ind]);
if i == n - 1 || label.abs() < labels[order[i + 1]].abs() {
exp_p_sum -= accumulated_sum;
accumulated_sum = 0.0;
}
}
out / num_events as f64
}
fn supports_label_matrix(&self) -> bool {
false
}
}
const PARALLEL_INTERVAL_ROWS: usize = 16_384;
fn interval_mean(preds: &[f32], info: &MetaInfo, row: impl Fn(f64, f64, f64) -> f64 + Sync) -> f64 {
let (lower, upper) = match info.bounds {
Some(bounds) => (bounds.lower(), bounds.upper()),
None => (info.label_values(), info.label_values()),
};
if upper.len() != lower.len() || !consistent(preds, lower, info.weights, 1) {
return f64::NAN;
}
let value = |i: usize| {
row(
f64::from(lower[i]),
f64::from(upper[i]),
f64::from(preds[i]),
)
};
let weight = |i: usize| info.weights.map_or(1.0, |w| f64::from(w[i]));
let n = lower.len();
let mut residue_sum = 0.0f64;
let mut weights_sum = 0.0f64;
if n >= PARALLEL_INTERVAL_ROWS && rayon::current_num_threads() > 1 {
let values: Vec<f64> = (0..n)
.into_par_iter()
.with_min_len(4096)
.map(value)
.collect();
for (i, v) in values.into_iter().enumerate() {
let w = weight(i);
residue_sum += v * w;
weights_sum += w;
}
} else {
for (i, ((&lo, &hi), &pred)) in lower.iter().zip(upper).zip(preds).enumerate() {
let w = weight(i);
residue_sum += row(f64::from(lo), f64::from(hi), f64::from(pred)) * w;
weights_sum += w;
}
}
if weights_sum == 0.0 {
residue_sum
} else {
residue_sum / weights_sum
}
}
fn validate_intervals(name: &str, info: &MetaInfo) -> Result<()> {
if info.n_rows > 0 && info.bounds.is_none() && info.label_values().is_empty() {
return Err(HessboostError::invalid_data(
"label_bounds",
format!("missing; metric `{name}` needs label bounds or labels"),
));
}
Ok(())
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct AftNLogLik {
distribution: AftDistribution,
sigma: f32,
}
impl AftNLogLik {
pub(crate) fn new(distribution: AftDistribution, sigma: f32) -> Self {
AftNLogLik {
distribution,
sigma,
}
}
}
impl Metric for AftNLogLik {
fn name(&self) -> &'static str {
"aft-nloglik"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_info(preds, &MetaInfo::new(labels, weights, None))
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
let sigma = f64::from(self.sigma);
interval_mean(preds, info, |lo, hi, pred| {
aft_nloglik(self.distribution, lo, hi, pred, sigma)
})
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
validate_intervals(self.name(), info)
}
fn supports_label_matrix(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) struct IntervalRegressionAccuracy;
impl Metric for IntervalRegressionAccuracy {
fn name(&self) -> &'static str {
"interval-regression-accuracy"
}
fn maximize(&self) -> bool {
true
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_info(preds, &MetaInfo::new(labels, weights, None))
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
interval_mean(preds, info, |lo, hi, log_pred| {
let pred = log_pred.exp();
if pred >= lo && pred <= hi { 1.0 } else { 0.0 }
})
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
validate_intervals(self.name(), info)
}
fn supports_label_matrix(&self) -> bool {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
fn bounded<'a>(lower: &'a [f32], upper: &'a [f32], weights: Option<&'a [f32]>) -> MetaInfo<'a> {
MetaInfo {
n_rows: lower.len(),
bounds: Some(crate::data::LabelBounds::new(lower, upper)),
weights,
..MetaInfo::unlabeled(0)
}
}
#[test]
fn cox_nloglik_matches_breslow_likelihood() {
let labels = [2.0f32, 1.0, -2.0, 3.0];
let h = [2.0f32, 1.0, 0.5, 4.0];
let total = 7.5f64;
let want = (-(1.0f64.ln() - total.ln())
- (2.0f64.ln() - (total - 1.0).ln())
- (4.0f64.ln() - 4.0f64.ln()))
/ 3.0;
assert_relative_eq!(
CoxNLogLik.eval(&h, &labels, None),
want,
max_relative = 1e-7
);
}
#[test]
fn cox_nloglik_refuses_label_matrix() {
use crate::prelude::{DMatrix, TrainingParams, train};
let x: Vec<f32> = (0..4).map(|i| i as f32).collect();
let two_targets = DMatrix::from_dense(&x, 4, 1)
.unwrap()
.with_label_matrix(&[1.0, 2.0, -3.0, 4.0, 2.0, 1.0, 3.0, -4.0], 2)
.unwrap();
let params = TrainingParams::builder()
.eval_metric(crate::metric::EvalMetric::CoxNLogLik)
.build()
.unwrap();
assert!(matches!(
train(¶ms, &two_targets, 1),
Err(HessboostError::InvalidParameter { name, .. }) if name == "eval_metric"
));
}
#[test]
fn interval_accuracy_counts_inclusive_hits_weighted() {
let lower = [1.0f32, 0.0, 1.5, 5.0];
let upper = [1.0f32, 3.0, f32::INFINITY, 6.0];
let margins = [0.0f32, 2.0f32.ln(), 2.0f32.ln(), 0.0];
let weights = [1.0f32, 1.0, 1.0, 3.0];
let info = bounded(&lower, &upper, Some(&weights));
let m = IntervalRegressionAccuracy;
assert!(m.maximize());
assert_relative_eq!(m.eval_info(&margins, &info), 3.0 / 6.0);
}
#[test]
fn aft_nloglik_uncensored_normal_is_lognormal_density() {
let lower = [1.5f32, 0.7];
let margins = [0.2f32, -0.1];
let sigma = 0.5f64;
let info = bounded(&lower, &lower, None);
let want: f64 = lower
.iter()
.zip(&margins)
.map(|(&t, &m)| {
let t = f64::from(t);
let z = (t.ln() - f64::from(m)) / sigma;
0.5 * z * z + (sigma * t * (2.0 * std::f64::consts::PI).sqrt()).ln()
})
.sum::<f64>()
/ 2.0;
let m = AftNLogLik::new(AftDistribution::Normal, 0.5);
assert_relative_eq!(m.eval_info(&margins, &info), want, max_relative = 1e-12);
assert_relative_eq!(m.eval(&margins, &lower, None), want, max_relative = 1e-12);
}
#[test]
fn default_aft_metric_uses_unit_scale() {
use crate::metric::{DEFAULT_SOURCE, XgboostMetricSource, named};
use crate::objective::{AftLoss, Loss};
let lower = [1.5f32, 0.0, 2.0];
let upper = [1.5f32, 3.0, f32::INFINITY];
let margins = [0.2f32, 0.5, 0.1];
let info = bounded(&lower, &upper, None);
let default = AftLoss::new(AftDistribution::Logistic, 0.8)
.default_metric()
.build(1)
.unwrap();
let source = XgboostMetricSource {
aft_loss_distribution: AftDistribution::Logistic,
aft_loss_distribution_scale: 0.8,
..DEFAULT_SOURCE
};
let explicit = named("aft-nloglik", 1, &source).unwrap();
let unit = AftNLogLik::new(AftDistribution::Logistic, 1.0).eval_info(&margins, &info);
let scaled = AftNLogLik::new(AftDistribution::Logistic, 0.8).eval_info(&margins, &info);
assert_ne!(unit, scaled);
assert_eq!(default.eval_info(&margins, &info), unit);
assert_eq!(explicit.eval_info(&margins, &info), scaled);
}
}