use serde::{Deserialize, Serialize};
use crate::error::{HessboostError, Result};
use std::num::NonZeroUsize;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum FeatureType {
#[default]
Numerical,
Categorical,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
#[non_exhaustive]
pub struct GroupInfo {
pub group_ptr: Vec<usize>,
}
impl GroupInfo {
pub fn from_sizes(sizes: &[usize]) -> Self {
let mut group_ptr = Vec::with_capacity(sizes.len() + 1);
group_ptr.push(0);
let mut acc = 0usize;
for &s in sizes {
acc = acc.saturating_add(s);
group_ptr.push(acc);
}
GroupInfo { group_ptr }
}
pub fn num_groups(&self) -> usize {
self.group_ptr.len().saturating_sub(1)
}
pub fn num_rows(&self) -> usize {
self.group_ptr.last().copied().unwrap_or(0)
}
pub fn iter_ranges(&self) -> impl Iterator<Item = (usize, usize)> + '_ {
self.group_ptr.windows(2).map(|w| (w[0], w[1]))
}
pub(crate) fn partitions(&self, n_rows: usize) -> bool {
self.group_ptr.first() == Some(&0)
&& self.group_ptr.last() == Some(&n_rows)
&& self.group_ptr.is_sorted()
}
}
#[derive(Debug, Clone, Copy)]
pub struct Labels<'a> {
values: &'a [f32],
n_targets: NonZeroUsize,
}
impl<'a> Labels<'a> {
pub fn new(values: &'a [f32], n_targets: NonZeroUsize) -> Self {
Labels { values, n_targets }
}
pub fn single(values: &'a [f32]) -> Self {
Labels::new(values, NonZeroUsize::MIN)
}
#[inline]
pub fn values(&self) -> &'a [f32] {
self.values
}
#[inline]
pub fn n_targets(&self) -> NonZeroUsize {
self.n_targets
}
}
#[derive(Debug, Clone, Copy)]
pub struct LabelBounds<'a> {
lower: &'a [f32],
upper: &'a [f32],
}
impl<'a> LabelBounds<'a> {
pub fn new(lower: &'a [f32], upper: &'a [f32]) -> Self {
LabelBounds { lower, upper }
}
#[inline]
pub fn lower(&self) -> &'a [f32] {
self.lower
}
#[inline]
pub fn upper(&self) -> &'a [f32] {
self.upper
}
}
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub struct MetaInfo<'a> {
pub n_rows: usize,
pub labels: Option<Labels<'a>>,
pub weights: Option<&'a [f32]>,
pub group: Option<&'a GroupInfo>,
pub bounds: Option<LabelBounds<'a>>,
}
impl<'a> MetaInfo<'a> {
pub fn new(
labels: &'a [f32],
weights: Option<&'a [f32]>,
group: Option<&'a GroupInfo>,
) -> Self {
MetaInfo {
n_rows: labels.len(),
labels: Some(Labels::single(labels)),
weights,
group,
bounds: None,
}
}
pub fn unlabeled(n_rows: usize) -> Self {
MetaInfo {
n_rows,
labels: None,
weights: None,
group: None,
bounds: None,
}
}
#[inline]
pub fn n_targets(&self) -> usize {
self.labels.map_or(1, |labels| labels.n_targets().get())
}
#[inline]
pub fn label_values(&self) -> &'a [f32] {
self.labels.map_or(&[], |labels| labels.values())
}
pub(crate) fn check_layout(&self) -> Result<()> {
let n_targets = self.n_targets();
let cells = self.n_rows.checked_mul(n_targets);
if cells.is_none() {
return Err(HessboostError::invalid_data(
"labels",
format!(
"{n_targets} targets for {} rows overflow usize",
self.n_rows
),
));
}
if let Some(labels) = self.labels
&& cells != Some(labels.values().len())
{
return Err(self.label_count_error());
}
for (name, values) in [
("weights", self.weights),
("label_lower_bound", self.bounds.map(|b| b.lower())),
("label_upper_bound", self.bounds.map(|b| b.upper())),
] {
if let Some(values) = values
&& values.len() != self.n_rows
{
return Err(HessboostError::invalid_data(
name,
format!("{} {name} for {} rows", values.len(), self.n_rows),
));
}
}
Ok(())
}
fn label_count_error(&self) -> HessboostError {
HessboostError::invalid_data(
"labels",
format!(
"{} labels for {} rows of {} targets",
self.label_values().len(),
self.n_rows,
self.n_targets()
),
)
}
pub(crate) fn cell_weights(&self) -> Result<Option<Vec<f32>>> {
self.check_layout()?;
if self.labels.is_none() {
return Err(self.label_count_error());
}
let n_targets = self.n_targets();
Ok(self.weights.map(|w| {
w.iter()
.flat_map(|&wi| std::iter::repeat_n(wi, n_targets))
.collect()
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn group_prefix_sum() {
let g = GroupInfo::from_sizes(&[3, 2, 4]);
assert_eq!(g.group_ptr, vec![0, 3, 5, 9]);
assert_eq!(g.num_groups(), 3);
assert_eq!(g.num_rows(), 9);
let ranges: Vec<_> = g.iter_ranges().collect();
assert_eq!(ranges, vec![(0, 3), (3, 5), (5, 9)]);
}
#[test]
fn cell_weights_need_a_consistent_layout() {
let labels = [1.0, 2.0, 3.0, 4.0];
let weights = [1.0, 2.0];
let matrix = |k: usize| Some(Labels::new(&labels, NonZeroUsize::new(k).unwrap()));
let info = MetaInfo {
n_rows: 2,
labels: matrix(2),
..MetaInfo::new(&labels, Some(&weights), None)
};
assert_eq!(info.cell_weights().unwrap(), Some(vec![1.0, 1.0, 2.0, 2.0]));
for bad in [
MetaInfo {
labels: matrix(usize::MAX),
..info
},
MetaInfo {
labels: matrix(3),
..info
},
MetaInfo {
labels: None,
..info
},
MetaInfo {
weights: Some(&weights[..1]),
..info
},
MetaInfo {
bounds: Some(LabelBounds::new(&labels, &weights)),
..info
},
] {
assert!(bad.cell_weights().is_err(), "{bad:?}");
}
}
#[test]
fn partitions_needs_ordered_ptr_spanning_the_rows() {
assert!(GroupInfo::from_sizes(&[2, 0, 1]).partitions(3));
assert!(!GroupInfo::from_sizes(&[2]).partitions(3));
assert!(!GroupInfo::default().partitions(0));
let unordered = GroupInfo {
group_ptr: vec![0, 3, 1, 3],
};
assert!(!unordered.partitions(3));
assert!(!GroupInfo::from_sizes(&[usize::MAX, 2]).partitions(1));
}
}