use super::{GradPair, MIN_HESS, Objective};
#[derive(Debug, Clone, Copy)]
pub struct SoftmaxObjective {
num_class: usize,
output_prob: bool,
}
impl SoftmaxObjective {
pub fn new(num_class: usize, output_prob: bool) -> Self {
SoftmaxObjective {
num_class,
output_prob,
}
}
}
impl Objective for SoftmaxObjective {
fn name(&self) -> &str {
if self.output_prob {
"multi:softprob"
} else {
"multi:softmax"
}
}
fn n_outputs(&self) -> usize {
self.num_class
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let k = self.num_class;
let n = labels.len();
super::check_gradient_inputs(n, k, preds, labels, weights, out);
super::rowwise_gradient(n, k, preds, labels, weights, out, |p, l, w, o| {
crate::simd::softmax_gradient(p, l, w, k, MIN_HESS, o);
});
}
fn pred_transform(&self, preds: &mut [f32]) {
let k = self.num_class;
crate::simd::softmax_rows_inplace(preds, k);
}
fn base_margins(
&self,
labels: &[f32],
weights: Option<&[f32]>,
_group: Option<&crate::data::GroupInfo>,
) -> Vec<f32> {
let k = self.num_class;
let mut margins = vec![0.0f32; k];
for (i, &y) in labels.iter().enumerate() {
if let Some(slot) = margins.get_mut(y as usize) {
*slot += weights.map_or(1.0, |ws| ws[i]);
}
}
let sum_w = match weights {
Some(ws) => ws.iter().map(|&w| f64::from(w)).sum::<f64>(),
None => f64::from(labels.len() as f32),
};
let inv_sum_w = 1.0 / sum_w;
for m in &mut margins {
*m = ((f64::from(*m) * inv_sum_w) as f32 + 1e-6).ln();
}
let n = k as f32;
let mean = margins.iter().map(|m| m / n).sum::<f32>();
for m in &mut margins {
*m -= mean;
}
margins
}
fn default_metric(&self) -> String {
"mlogloss".to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn softmax_normalizes() {
let mut r = [1.0f32, 2.0, 3.0];
let num_class = r.len();
crate::simd::softmax_rows_inplace(&mut r, num_class);
assert_relative_eq!(r.iter().sum::<f32>(), 1.0, epsilon = 1e-6);
assert!(r[2] > r[1] && r[1] > r[0]);
}
#[test]
fn gradient_layout_and_values() {
let obj = SoftmaxObjective::new(3, true);
let preds = [0.0f32; 6];
let labels = [0.0f32, 2.0];
let mut out = vec![GradPair::default(); 6];
obj.gradient(&preds, &labels, None, &mut out);
assert_relative_eq!(out[0].grad, 1.0 / 3.0 - 1.0, epsilon = 1e-6);
assert_relative_eq!(out[1].grad, 1.0 / 3.0, epsilon = 1e-6);
assert_relative_eq!(out[5].grad, 1.0 / 3.0 - 1.0, epsilon = 1e-6);
assert_relative_eq!(out[0].hess, 2.0 * (1.0 / 3.0) * (2.0 / 3.0), epsilon = 1e-6);
}
#[test]
fn base_margins_are_centered_log_frequencies() {
let obj = SoftmaxObjective::new(3, true);
let labels = [0.0f32, 0.0, 1.0, 2.0];
let m = obj.base_margins(&labels, None, None);
assert_eq!(m.len(), 3);
assert!(m.iter().sum::<f32>().abs() < 1e-6);
let expected_gap = (0.5f32 + 1e-6).ln() - (0.25f32 + 1e-6).ln();
assert!((m[0] - m[1] - expected_gap).abs() < 1e-6, "{m:?}");
assert_eq!(m[1], m[2]);
let w = [1.0f32, 1.0, 2.0, 2.0];
let uniform = obj.base_margins(&labels, Some(&w), None);
assert!(uniform.iter().all(|v| v.abs() < 1e-6), "{uniform:?}");
}
}