use super::{GradPair, Objective};
use crate::data::GroupInfo;
use crate::metric::{argsort_desc, group_ranges};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RankMode {
Pairwise,
Ndcg,
Map,
}
#[derive(Debug, Clone, Copy)]
pub struct LambdaMartObjective {
mode: RankMode,
top_k: usize,
}
impl LambdaMartObjective {
pub fn pairwise(top_k: usize) -> Self {
LambdaMartObjective {
mode: RankMode::Pairwise,
top_k,
}
}
pub fn ndcg(top_k: usize) -> Self {
LambdaMartObjective {
mode: RankMode::Ndcg,
top_k,
}
}
pub fn map(top_k: usize) -> Self {
LambdaMartObjective {
mode: RankMode::Map,
top_k,
}
}
#[allow(clippy::too_many_arguments)]
fn accumulate_group(
&self,
preds: &[f32],
labels: &[f32],
start: usize,
end: usize,
query_weight: f32,
weight_norm: f32,
out: &mut [GradPair],
) {
let n = end - start;
if n < 2 {
return;
}
let p = &preds[start..end];
let y = &labels[start..end];
let order = argsort_desc(p);
let metric = MetricCtx::build(self.mode, y, &order, self.top_k);
let best_score = p[order[0]];
let worst_score = p[*order.last().unwrap()];
let mut sum_lambda = 0.0f64;
for i in 0..n.min(self.top_k) {
for j in i + 1..n {
let mut rank_high = i;
let mut rank_low = j;
let mut idx_high = order[rank_high];
let mut idx_low = order[rank_low];
if y[idx_high] == y[idx_low] {
continue;
}
if y[idx_high] < y[idx_low] {
std::mem::swap(&mut rank_high, &mut rank_low);
std::mem::swap(&mut idx_high, &mut idx_low);
}
let score_diff = p[idx_high] - p[idx_low]; let delta_score = f64::from(score_diff.abs());
let sigmoid = f64::from(1.0f32 / ((-score_diff).min(88.7).exp() + 1.0));
let mut delta = metric
.delta(y[idx_high], y[idx_low], rank_high, rank_low)
.abs();
if best_score != worst_score {
delta /= delta_score + 0.01;
}
let lambda = (sigmoid - 1.0) * delta;
let hessian = (sigmoid * (1.0 - sigmoid)).max(1e-16) * delta * 2.0;
let pg = GradPair::new(lambda as f32, hessian as f32);
out[start + idx_high].grad += pg.grad;
out[start + idx_high].hess += pg.hess;
out[start + idx_low].grad -= pg.grad;
out[start + idx_low].hess += pg.hess;
sum_lambda += -2.0 * f64::from(pg.grad);
}
}
let norm = if sum_lambda > 0.0 {
(sum_lambda + 1.0).log2() / sum_lambda
} else {
1.0
};
let group = &mut out[start..end];
if norm != 1.0 {
let norm = norm as f32;
for g in group.iter_mut() {
g.grad *= norm;
g.hess *= norm;
}
}
for g in group.iter_mut() {
g.grad *= query_weight;
g.hess *= query_weight;
g.grad *= weight_norm;
g.hess *= weight_norm;
}
}
fn compute(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
group: Option<&GroupInfo>,
out: &mut [GradPair],
) {
super::check_gradient_inputs(labels.len(), 1, preds, labels, weights, out);
out.fill(GradPair::default());
let ranges = group_ranges(preds.len(), group);
let group_weights: Vec<f32> = ranges
.iter()
.map(|&(start, _)| weights.map_or(1.0, |w| w[start]))
.collect();
let sum_w: f64 = group_weights.iter().map(|&w| f64::from(w)).sum();
let weight_norm = if sum_w == 0.0 {
0.0
} else {
(ranges.len() as f64 / sum_w) as f32
};
for ((start, end), weight) in ranges.into_iter().zip(group_weights) {
self.accumulate_group(preds, labels, start, end, weight, weight_norm, out);
}
}
}
impl Objective for LambdaMartObjective {
fn name(&self) -> &str {
match self.mode {
RankMode::Pairwise => "rank:pairwise",
RankMode::Ndcg => "rank:ndcg",
RankMode::Map => "rank:map",
}
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
self.compute(preds, labels, weights, None, out);
}
fn gradient_grouped(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
group: Option<&GroupInfo>,
out: &mut [GradPair],
) {
self.compute(preds, labels, weights, group, out);
}
fn default_metric(&self) -> String {
let base = match self.mode {
RankMode::Pairwise | RankMode::Ndcg => "ndcg",
RankMode::Map => "map",
};
format!("{base}@{}", self.top_k)
}
}
enum MetricCtx {
Uniform,
Ndcg { discounts: Vec<f64>, inv_idcg: f64 },
Map { n_rel: Vec<f64>, acc: Vec<f64> },
}
impl MetricCtx {
fn build(mode: RankMode, labels: &[f32], order: &[usize], top_k: usize) -> Self {
match mode {
RankMode::Pairwise => MetricCtx::Uniform,
RankMode::Ndcg => {
let discounts: Vec<f64> = (0..labels.len())
.map(|i| 1.0 / ((i + 2) as f64).log2())
.collect();
let mut ideal: Vec<usize> = (0..labels.len()).collect();
ideal.sort_by(|&a, &b| labels[b].total_cmp(&labels[a])); let idcg: f64 = ideal
.iter()
.take(labels.len().min(top_k))
.enumerate()
.map(|(rank, &idx)| {
let gain = f64::from((1u32 << labels[idx] as u32) - 1);
discounts[rank] * gain
})
.sum();
MetricCtx::Ndcg {
discounts,
inv_idcg: if idcg == 0.0 { 0.0 } else { 1.0 / idcg },
}
}
RankMode::Map => {
let mut n_rel = vec![0.0; labels.len()];
let mut acc = vec![0.0; labels.len()];
for (rank, &idx) in order.iter().enumerate() {
let y = f64::from(labels[idx]);
n_rel[rank] = y + if rank == 0 { 0.0 } else { n_rel[rank - 1] };
acc[rank] = y / (rank + 1) as f64 + if rank == 0 { 0.0 } else { acc[rank - 1] };
}
MetricCtx::Map { n_rel, acc }
}
}
}
fn delta(&self, y_high: f32, y_low: f32, rank_high: usize, rank_low: usize) -> f64 {
match self {
MetricCtx::Uniform => 1.0,
MetricCtx::Ndcg {
discounts,
inv_idcg,
} => {
let gain_high = f64::from((1u32 << y_high as u32) - 1);
let gain_low = f64::from((1u32 << y_low as u32) - 1);
let original = gain_high * discounts[rank_high] + gain_low * discounts[rank_low];
let changed = gain_low * discounts[rank_high] + gain_high * discounts[rank_low];
(original - changed) * inv_idcg
}
MetricCtx::Map { n_rel, acc } => {
let (mut rh, mut rl, mut yh, mut yl) =
(rank_high, rank_low, f64::from(y_high), f64::from(y_low));
if rh > rl {
std::mem::swap(&mut rh, &mut rl);
std::mem::swap(&mut yh, &mut yl);
}
let total = *n_rel.last().unwrap();
let (m, n) = (n_rel[rl], n_rel[rh]);
let b = acc[rl - 1] - acc[rh];
if yh < yl {
(m / (rl + 1) as f64 - (n + 1.0) / (rh + 1) as f64 - b) / total
} else {
(n / (rh + 1) as f64 - m / (rl + 1) as f64 + b) / total
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::data::GroupInfo;
#[test]
fn pairwise_pushes_relevant_up() {
let obj = LambdaMartObjective::ndcg(32);
let preds = [0.0f32, 0.0, 0.0];
let labels = [2.0f32, 1.0, 0.0];
let g = GroupInfo::from_sizes(&[3]);
let mut out = vec![GradPair::default(); 3];
obj.gradient_grouped(&preds, &labels, None, Some(&g), &mut out);
assert!(out[0].grad < out[1].grad, "{out:?}");
assert!(out[1].grad < out[2].grad, "{out:?}");
assert!(out[0].grad < 0.0 && out[2].grad > 0.0);
assert!(out.iter().all(|g| g.hess >= 0.0));
}
#[test]
fn no_pairs_when_all_labels_equal() {
let obj = LambdaMartObjective::pairwise(32);
let preds = [0.5f32, -0.2, 1.0];
let labels = [1.0f32, 1.0, 1.0];
let g = GroupInfo::from_sizes(&[3]);
let mut out = vec![GradPair::default(); 3];
obj.gradient_grouped(&preds, &labels, None, Some(&g), &mut out);
assert!(out.iter().all(|g| g.grad == 0.0 && g.hess == 0.0));
}
#[test]
fn groups_are_independent() {
let obj = LambdaMartObjective::pairwise(32);
let preds = [0.0f32, 0.0, 0.0, 0.0];
let labels = [1.0f32, 0.0, 0.0, 1.0];
let g = GroupInfo::from_sizes(&[2, 2]);
let mut out = vec![GradPair::default(); 4];
obj.gradient_grouped(&preds, &labels, None, Some(&g), &mut out);
assert!(out[0].grad < 0.0 && out[1].grad > 0.0);
assert!(out[3].grad < 0.0 && out[2].grad > 0.0);
}
fn pairwise_pair(s_high: f32, s_low: f32) -> (f32, f32) {
let diff = s_high - s_low;
let sigmoid = f64::from(1.0f32 / ((-diff).min(88.7).exp() + 1.0));
let delta = 1.0 / (f64::from(diff.abs()) + 0.01);
let lambda = (sigmoid - 1.0) * delta;
let hessian = (sigmoid * (1.0 - sigmoid)).max(1e-16) * delta * 2.0;
(lambda as f32, hessian as f32)
}
#[test]
fn signed_zero_scores_tie_in_input_order() {
let obj = LambdaMartObjective::pairwise(1);
let preds = [-0.0f32, 0.0, -1.0];
let labels = [2.0f32, 1.0, 0.0];
let g = GroupInfo::from_sizes(&[3]);
let mut out = vec![GradPair::default(); 3];
obj.gradient_grouped(&preds, &labels, None, Some(&g), &mut out);
let (g01, h01) = pairwise_pair(preds[0], preds[1]);
let (g02, h02) = pairwise_pair(preds[0], preds[2]);
let sum_lambda = -2.0 * f64::from(g01) + -2.0 * f64::from(g02);
let norm = ((sum_lambda + 1.0).log2() / sum_lambda) as f32;
assert_ne!(out[0].grad, 0.0);
assert_eq!(out[0].grad, (g01 + g02) * norm, "{out:?}");
assert_eq!(out[0].hess, (h01 + h02) * norm, "{out:?}");
assert_eq!(out[1].grad, -g01 * norm, "{out:?}");
assert_eq!(out[1].hess, h01 * norm, "{out:?}");
assert_eq!(out[2].grad, -g02 * norm, "{out:?}");
assert_eq!(out[2].hess, h02 * norm, "{out:?}");
}
#[test]
fn query_weights_scale_as_sequential_f32_products() {
let obj = LambdaMartObjective::ndcg(32);
let preds = [0.3f32, -0.7, 1.1, 0.2, -0.4, 0.9, 0.05];
let labels = [2.0f32, 0.0, 1.0, 3.0, 1.0, 0.0, 2.0];
let g = GroupInfo::from_sizes(&[3, 4]);
let group_w = [0.3f32, 1.7];
let weights = [0.3f32, 0.3, 0.3, 1.7, 1.7, 1.7, 1.7];
let mut normed = vec![GradPair::default(); 7];
obj.gradient_grouped(&preds, &labels, None, Some(&g), &mut normed);
let mut out = vec![GradPair::default(); 7];
obj.gradient_grouped(&preds, &labels, Some(&weights), Some(&g), &mut out);
let sum_w = f64::from(group_w[0]) + f64::from(group_w[1]);
let w_norm = (2.0 / sum_w) as f32;
for (i, (got, base)) in out.iter().zip(&normed).enumerate() {
let w = if i < 3 { group_w[0] } else { group_w[1] };
assert_eq!(
got.grad.to_bits(),
((base.grad * w) * w_norm).to_bits(),
"{i}"
);
assert_eq!(
got.hess.to_bits(),
((base.hess * w) * w_norm).to_bits(),
"{i}"
);
}
}
#[test]
fn default_metric_follows_mode_and_top_k() {
assert_eq!(
LambdaMartObjective::pairwise(32).default_metric(),
"ndcg@32"
);
assert_eq!(LambdaMartObjective::ndcg(5).default_metric(), "ndcg@5");
assert_eq!(LambdaMartObjective::map(10).default_metric(), "map@10");
}
}