use super::{GradPair, check_label_domain, check_label_width};
use crate::data::MetaInfo;
use crate::error::{HessboostError, Result};
use rayon::prelude::*;
pub(super) const PARALLEL_QUERY_ROWS: usize = 4096;
pub(super) fn validate_query_info(info: &MetaInfo, invalid: impl Fn(f32) -> bool) -> Result<()> {
info.check_layout()?;
check_label_width(info, 1)?;
check_label_domain(info, invalid)?;
let Some(group) = info.group else {
return Err(HessboostError::invalid_data(
"group_sizes",
"ranking objectives require query groups",
));
};
if !group.partitions(info.n_rows) || group.iter_ranges().any(|(start, end)| start == end) {
return Err(HessboostError::invalid_data(
"group_sizes",
format!(
"query groups are not non-empty consecutive row ranges covering the {} rows",
info.n_rows
),
));
}
if let Some(weights) = info.weights {
for (start, end) in group.iter_ranges() {
if weights[start..end]
.iter()
.any(|weight| *weight != weights[start])
{
return Err(HessboostError::invalid_data(
"weights",
"ranking objectives require one constant weight per query group",
));
}
}
}
Ok(())
}
pub(super) fn for_each_query<S>(
ranges: &[(usize, usize)],
out: &mut [GradPair],
init: impl Fn() -> S + Sync + Send,
f: impl Fn(&mut S, usize, usize, &mut [GradPair]) + Sync + Send,
) {
let n_rows = out.len();
let mut queries = Vec::with_capacity(ranges.len());
let mut rest = out;
let mut offset = 0;
for &(start, end) in ranges {
let (_, tail) = std::mem::take(&mut rest).split_at_mut(start - offset);
let (rows, tail) = tail.split_at_mut(end - start);
rest = tail;
offset = end;
queries.push((start, rows));
}
let process = |scratch: &mut S, (query, (start, rows)): (usize, (usize, &mut [GradPair]))| {
f(scratch, query, start, rows);
};
if n_rows >= PARALLEL_QUERY_ROWS && queries.len() > 1 && rayon::current_num_threads() > 1 {
queries
.into_par_iter()
.enumerate()
.for_each_init(&init, process);
} else {
let mut scratch = init();
for item in queries.into_iter().enumerate() {
process(&mut scratch, item);
}
}
}