use crate::data::meta::{FeatureType, GroupInfo, LabelBounds, Labels, MetaInfo};
use crate::error::{HessboostError, Result};
use rayon::prelude::*;
use std::num::NonZeroUsize;
const PARALLEL_COPY_VALUES: usize = 1 << 20;
#[inline]
pub(crate) fn is_missing(v: f32, missing: f32) -> bool {
if missing.is_nan() {
v.is_nan()
} else {
v == missing
}
}
#[inline]
pub(crate) fn check_len(what: &'static str, got: usize, expected: usize) -> Result<()> {
if got != expected {
return Err(HessboostError::dimension_mismatch(what, expected, got));
}
Ok(())
}
fn check_csr(indptr: &[usize], nnz: usize) -> Result<()> {
if indptr[0] != 0 {
return Err(HessboostError::invalid_data(
"csr indptr",
"the first offset must be 0",
));
}
for pair in indptr.windows(2) {
if pair[0] > pair[1] || pair[1] > nnz {
return Err(HessboostError::invalid_data(
"csr indptr",
"offsets must be monotonic and within the values array",
));
}
}
check_len("csr indptr terminal", indptr[indptr.len() - 1], nnz)
}
fn check_finite(name: &'static str, values: &[f32], reason: &'static str) -> Result<()> {
if values.iter().any(|v| !v.is_finite()) {
return Err(HessboostError::invalid_data(name, reason));
}
Ok(())
}
fn check_weights(
name: &'static str,
weights: &[f32],
invalid: &'static str,
none_positive: &'static str,
) -> Result<()> {
if weights.iter().any(|v| !v.is_finite() || *v < 0.0) {
return Err(HessboostError::invalid_data(name, invalid));
}
if !weights.iter().any(|v| *v > 0.0) {
return Err(HessboostError::invalid_data(name, none_positive));
}
Ok(())
}
fn check_dense(data: &[f32], n_rows: usize, n_cols: usize, missing: f32) -> Result<bool> {
if n_rows == 0 || n_cols == 0 {
return Err(HessboostError::EmptyDataset(
"from_dense: zero rows or columns",
));
}
let expected = n_rows
.checked_mul(n_cols)
.ok_or_else(|| HessboostError::invalid_data("data", "n_rows * n_cols overflows usize"))?;
check_len("dense data length", data.len(), expected)?;
let parallel = data.len() >= PARALLEL_COPY_VALUES && rayon::current_num_threads() > 1;
let rejects = |chunk: &[f32]| {
if missing.is_nan() {
chunk.iter().any(|v| v.is_infinite())
} else {
chunk.iter().any(|&v| v != missing && !v.is_finite())
}
};
let invalid = if parallel {
data.par_chunks(PARALLEL_COPY_VALUES / 16).any(rejects)
} else {
rejects(data)
};
if invalid {
return Err(HessboostError::invalid_data(
"data",
"non-missing feature values must be finite",
));
}
Ok(parallel)
}
#[derive(Debug, Clone)]
enum Storage {
Dense(Vec<f32>),
Csr {
indptr: Vec<usize>,
indices: Vec<u32>,
values: Vec<f32>,
},
}
#[derive(Debug, Clone)]
pub struct DMatrix {
n_rows: usize,
n_cols: usize,
storage: Storage,
missing: f32,
labels: Option<Vec<f32>>,
n_targets: NonZeroUsize,
label_bounds: Option<(Vec<f32>, Vec<f32>)>,
weights: Option<Vec<f32>>,
base_margin: Option<Vec<f32>>,
group: Option<GroupInfo>,
feature_types: Vec<FeatureType>,
feature_weights: Option<Vec<f32>>,
}
impl DMatrix {
fn new(n_rows: usize, n_cols: usize, storage: Storage, missing: f32) -> Self {
DMatrix {
n_rows,
n_cols,
storage,
missing,
labels: None,
n_targets: NonZeroUsize::MIN,
label_bounds: None,
weights: None,
base_margin: None,
group: None,
feature_types: vec![FeatureType::Numerical; n_cols],
feature_weights: None,
}
}
pub fn from_dense(data: &[f32], n_rows: usize, n_cols: usize) -> Result<Self> {
Self::from_dense_with_missing(data, n_rows, n_cols, f32::NAN)
}
pub fn from_dense_with_missing(
data: &[f32],
n_rows: usize,
n_cols: usize,
missing: f32,
) -> Result<Self> {
let parallel = check_dense(data, n_rows, n_cols, missing)?;
let values = if parallel {
data.par_iter().copied().collect()
} else {
data.to_vec()
};
Ok(Self::new(n_rows, n_cols, Storage::Dense(values), missing))
}
pub(crate) fn from_dense_vec(data: Vec<f32>, n_rows: usize, n_cols: usize) -> Result<Self> {
check_dense(&data, n_rows, n_cols, f32::NAN)?;
Ok(Self::new(n_rows, n_cols, Storage::Dense(data), f32::NAN))
}
pub fn from_csr(
indptr: Vec<usize>,
indices: Vec<u32>,
values: Vec<f32>,
n_cols: usize,
) -> Result<Self> {
if indptr.is_empty() {
return Err(HessboostError::EmptyDataset("from_csr: empty indptr"));
}
let n_rows = indptr.len() - 1;
if n_rows == 0 || n_cols == 0 {
return Err(HessboostError::EmptyDataset(
"from_csr: zero rows or columns",
));
}
check_len("csr indices/values length", values.len(), indices.len())?;
check_csr(&indptr, values.len())?;
if let Some(&m) = indices.iter().max()
&& (m as usize) >= n_cols
{
return Err(HessboostError::FeatureOutOfBounds {
index: m as usize,
num_features: n_cols,
});
}
check_finite(
"csr values",
&values,
"stored feature values must be finite",
)?;
let mut seen = std::collections::HashSet::new();
for row in 0..n_rows {
seen.clear();
for &col in &indices[indptr[row]..indptr[row + 1]] {
if !seen.insert(col) {
return Err(HessboostError::invalid_data(
"csr indices",
format!("duplicate column {col} in row {row}"),
));
}
}
}
Ok(Self::new(
n_rows,
n_cols,
Storage::Csr {
indptr,
indices,
values,
},
f32::NAN,
))
}
pub fn with_labels(self, labels: &[f32]) -> Result<Self> {
self.with_label_matrix(labels, 1)
}
pub fn with_label_matrix(mut self, labels: &[f32], n_targets: usize) -> Result<Self> {
let Some(n_targets) = NonZeroUsize::new(n_targets) else {
return Err(HessboostError::invalid_data(
"labels",
"n_targets must be at least 1",
));
};
let expected = self.n_rows.checked_mul(n_targets.get()).ok_or_else(|| {
HessboostError::invalid_data("labels", "n_rows * n_targets overflows usize")
})?;
check_len("labels", labels.len(), expected)?;
check_finite("labels", labels, "all labels must be finite")?;
self.labels = Some(labels.to_vec());
self.n_targets = n_targets;
Ok(self)
}
pub fn with_label_bounds(mut self, lower: &[f32], upper: &[f32]) -> Result<Self> {
check_len("label_lower_bound", lower.len(), self.n_rows)?;
check_len("label_upper_bound", upper.len(), self.n_rows)?;
for (name, bound) in [("label_lower_bound", lower), ("label_upper_bound", upper)] {
if bound.iter().any(|v| v.is_nan()) {
return Err(HessboostError::invalid_data(
name,
"label bounds must not be NaN",
));
}
}
self.label_bounds = Some((lower.to_vec(), upper.to_vec()));
Ok(self)
}
pub fn with_feature_weights(mut self, weights: &[f32]) -> Result<Self> {
check_len("feature_weights", weights.len(), self.n_cols)?;
check_weights(
"feature_weights",
weights,
"feature weights must be finite and non-negative",
"at least one feature weight must be positive",
)?;
self.feature_weights = Some(weights.to_vec());
Ok(self)
}
pub fn with_weights(mut self, weights: &[f32]) -> Result<Self> {
check_len("weights", weights.len(), self.n_rows)?;
check_weights(
"weights",
weights,
"weights must be finite and non-negative",
"at least one weight must be positive",
)?;
self.weights = Some(weights.to_vec());
Ok(self)
}
pub fn with_base_margin(mut self, base_margin: &[f32]) -> Result<Self> {
if base_margin.is_empty() || !base_margin.len().is_multiple_of(self.n_rows) {
return Err(HessboostError::invalid_data(
"base_margin",
"length must be a non-zero multiple of n_rows",
));
}
check_finite("base_margin", base_margin, "all margins must be finite")?;
self.base_margin = Some(base_margin.to_vec());
Ok(self)
}
pub fn with_group_sizes(mut self, sizes: &[usize]) -> Result<Self> {
if sizes.is_empty() || sizes.contains(&0) {
return Err(HessboostError::invalid_data(
"group_sizes",
"groups must be non-empty and every group must contain a row",
));
}
let total = sizes.iter().try_fold(0usize, |acc, &s| acc.checked_add(s));
let Some(total) = total else {
return Err(HessboostError::invalid_data(
"group_sizes",
"group-size sum overflows usize",
));
};
check_len("group sizes sum", total, self.n_rows)?;
self.group = Some(GroupInfo::from_sizes(sizes));
Ok(self)
}
pub fn with_group_weights(mut self, weights: &[f32]) -> Result<Self> {
let group = self.group.as_ref().ok_or_else(|| {
HessboostError::invalid_data("group_weights", "attach group sizes first")
})?;
check_len("group_weights length", weights.len(), group.num_groups())?;
let invalid = "weights must be finite and non-negative with at least one positive value";
check_weights("group_weights", weights, invalid, invalid)?;
let mut expanded = Vec::with_capacity(self.n_rows);
for ((start, end), &weight) in group.iter_ranges().zip(weights) {
expanded.extend(std::iter::repeat_n(weight, end - start));
}
self.weights = Some(expanded);
Ok(self)
}
pub fn with_feature_types(mut self, types: &[FeatureType]) -> Result<Self> {
check_len("feature_types length", types.len(), self.n_cols)?;
self.feature_types = types.to_vec();
let invalid = |col: usize, v: f32| {
HessboostError::invalid_data(
"categorical feature",
format!("feature {col} contains invalid category value {v}"),
)
};
let is_invalid = |v: f32| v < 0.0 || v.fract() != 0.0 || v >= u32::MAX as f32;
match &self.storage {
Storage::Dense(_) => {
for (col, ty) in types.iter().enumerate() {
if *ty == FeatureType::Categorical {
for row in 0..self.n_rows {
if let Some(v) = self.get(row, col)
&& is_invalid(v)
{
return Err(invalid(col, v));
}
}
}
}
}
Storage::Csr { .. } => {
let mut first: Option<(usize, usize, f32)> = None;
self.for_each_entry(|row, col, v| {
let col = col as usize;
if types[col] == FeatureType::Categorical
&& is_invalid(v)
&& first.is_none_or(|(c, r, _)| (col, row) < (c, r))
{
first = Some((col, row, v));
}
});
if let Some((col, _, v)) = first {
return Err(invalid(col, v));
}
}
}
Ok(self)
}
#[inline]
pub fn n_rows(&self) -> usize {
self.n_rows
}
#[inline]
pub fn n_cols(&self) -> usize {
self.n_cols
}
#[inline]
pub fn missing(&self) -> f32 {
self.missing
}
#[inline]
pub fn labels(&self) -> Option<&[f32]> {
self.labels.as_deref()
}
#[inline]
pub fn n_targets(&self) -> usize {
self.n_targets.get()
}
#[inline]
pub fn label_lower_bound(&self) -> Option<&[f32]> {
self.label_bounds
.as_ref()
.map(|(lower, _)| lower.as_slice())
}
#[inline]
pub fn label_upper_bound(&self) -> Option<&[f32]> {
self.label_bounds
.as_ref()
.map(|(_, upper)| upper.as_slice())
}
#[inline]
pub fn weights(&self) -> Option<&[f32]> {
self.weights.as_deref()
}
#[inline]
pub fn base_margin(&self) -> Option<&[f32]> {
self.base_margin.as_deref()
}
#[inline]
pub fn group(&self) -> Option<&GroupInfo> {
self.group.as_ref()
}
#[inline]
pub fn feature_types(&self) -> &[FeatureType] {
&self.feature_types
}
#[inline]
pub fn feature_weights(&self) -> Option<&[f32]> {
self.feature_weights.as_deref()
}
pub fn info(&self) -> MetaInfo<'_> {
MetaInfo {
n_rows: self.n_rows,
labels: self
.labels
.as_deref()
.map(|values| Labels::new(values, self.n_targets)),
weights: self.weights.as_deref(),
group: self.group.as_ref(),
bounds: self
.label_bounds
.as_ref()
.map(|(lower, upper)| LabelBounds::new(lower, upper)),
}
}
pub fn get(&self, row: usize, col: usize) -> Option<f32> {
if row >= self.n_rows || col >= self.n_cols {
return None;
}
let v = match &self.storage {
Storage::Dense(data) => data[row * self.n_cols + col],
Storage::Csr {
indptr,
indices,
values,
} => {
let k = (indptr[row]..indptr[row + 1]).find(|&k| indices[k] as usize == col)?;
values[k]
}
};
(!is_missing(v, self.missing)).then_some(v)
}
#[inline]
pub(crate) fn dense_values(&self) -> Option<&[f32]> {
match &self.storage {
Storage::Dense(data) => Some(data),
Storage::Csr { .. } => None,
}
}
#[inline]
pub(crate) fn csr_parts(&self) -> Option<(&[usize], &[u32], &[f32])> {
match &self.storage {
Storage::Dense(_) => None,
Storage::Csr {
indptr,
indices,
values,
} => Some((indptr, indices, values)),
}
}
#[inline]
pub(crate) fn dense_values_mut(&mut self) -> Option<&mut [f32]> {
match &mut self.storage {
Storage::Dense(data) => Some(data),
Storage::Csr { .. } => None,
}
}
pub(crate) fn map_values(&self, mut f: impl FnMut(usize, usize, f32) -> f32) -> Self {
let mut out = self.clone();
match &mut out.storage {
Storage::Dense(data) => {
for (row, values) in data.chunks_exact_mut(self.n_cols).enumerate() {
for (col, v) in values.iter_mut().enumerate() {
*v = if is_missing(*v, self.missing) {
f32::NAN
} else {
f(row, col, *v)
};
}
}
}
Storage::Csr {
indptr,
indices,
values,
} => {
for row in 0..self.n_rows {
for k in indptr[row]..indptr[row + 1] {
if !is_missing(values[k], self.missing) {
values[k] = f(row, indices[k] as usize, values[k]);
}
}
}
}
}
out.missing = f32::NAN;
out
}
#[inline]
pub(crate) fn for_row_entry(&self, row: usize, mut f: impl FnMut(u32, f32)) {
match &self.storage {
Storage::Dense(data) => {
let base = row * self.n_cols;
for c in 0..self.n_cols {
let v = data[base + c];
if !is_missing(v, self.missing) {
f(c as u32, v);
}
}
}
Storage::Csr {
indptr,
indices,
values,
} => {
let (s, e) = (indptr[row], indptr[row + 1]);
for k in s..e {
let v = values[k];
if !is_missing(v, self.missing) {
f(indices[k], v);
}
}
}
}
}
pub(crate) fn to_csc(&self) -> CscView {
let mut col_counts = vec![0usize; self.n_cols];
self.for_each_entry(|_row, col, _v| col_counts[col as usize] += 1);
let mut col_ptr = vec![0usize; self.n_cols + 1];
for (c, count) in col_counts.iter_mut().enumerate() {
col_ptr[c + 1] = col_ptr[c] + *count;
*count = col_ptr[c];
}
let nnz = col_ptr[self.n_cols];
let mut rows = vec![0u32; nnz];
let mut vals = vec![0f32; nnz];
let mut cursor = col_counts;
self.for_each_entry(|row, col, v| {
let c = col as usize;
let pos = cursor[c];
rows[pos] = row as u32;
vals[pos] = v;
cursor[c] = pos + 1;
});
CscView {
n_rows: self.n_rows,
n_cols: self.n_cols,
col_ptr,
rows,
vals,
}
}
pub fn select_rows(&self, rows: &[usize]) -> Result<Self> {
if let Some(&row) = rows.iter().find(|&&row| row >= self.n_rows) {
return Err(HessboostError::invalid_param(
"rows",
format!("row index {row} is out of bounds for {} rows", self.n_rows),
));
}
let group_sizes = self
.selected_group_sizes(rows)
.map_err(|reason| HessboostError::invalid_param("rows", format!("rows {reason}")))?;
let mut out = match &self.storage {
Storage::Dense(data) => {
if rows.is_empty() {
return Err(HessboostError::EmptyDataset(
"from_dense: zero rows or columns",
));
}
let n_cols = self.n_cols;
let mut selected = Vec::with_capacity(rows.len() * n_cols);
for &r in rows {
selected.extend_from_slice(&data[r * n_cols..(r + 1) * n_cols]);
}
if !self.missing.is_nan() {
for v in &mut selected {
if *v == self.missing {
*v = f32::NAN;
}
}
}
Self::new(rows.len(), n_cols, Storage::Dense(selected), f32::NAN)
}
Storage::Csr { .. } => {
let mut indptr = Vec::with_capacity(rows.len() + 1);
indptr.push(0usize);
let mut indices: Vec<u32> = Vec::new();
let mut values: Vec<f32> = Vec::new();
for &r in rows {
self.for_row_entry(r, |index, value| {
indices.push(index);
values.push(value);
});
indptr.push(values.len());
}
DMatrix::from_csr(indptr, indices, values, self.n_cols)?
}
};
out.feature_types.clone_from(&self.feature_types);
out.feature_weights.clone_from(&self.feature_weights);
out.n_targets = self.n_targets;
let gather = |v: &[f32], stride: usize| {
let mut selected = Vec::with_capacity(rows.len() * stride);
for &r in rows {
selected.extend_from_slice(&v[r * stride..(r + 1) * stride]);
}
selected
};
out.labels = self
.labels
.as_deref()
.map(|l| gather(l, self.n_targets.get()));
out.label_bounds = self
.label_bounds
.as_ref()
.map(|(lo, hi)| (gather(lo, 1), gather(hi, 1)));
out.weights = self.weights.as_deref().map(|w| gather(w, 1));
out.base_margin = self
.base_margin
.as_deref()
.map(|bm| gather(bm, bm.len() / self.n_rows));
if let Some(sizes) = group_sizes.filter(|sizes| !sizes.is_empty()) {
out.group = Some(GroupInfo::from_sizes(&sizes));
}
Ok(out)
}
pub(crate) fn selected_group_sizes(
&self,
rows: &[usize],
) -> std::result::Result<Option<Vec<usize>>, String> {
let Some(group) = &self.group else {
return Ok(None);
};
let ptr = &group.group_ptr;
let mut sizes = Vec::new();
let mut at = 0;
while let Some(&row) = rows.get(at) {
let g = ptr.partition_point(|&start| start <= row) - 1;
let (start, end) = (ptr[g], ptr[g + 1]);
let run = rows.get(at..at + (end - start));
if !run.is_some_and(|run| run.iter().copied().eq(start..end)) {
return Err(format!(
"split query group {g} (rows {start}..{end}), reorder it, or interleave it \
with other rows; select whole groups, each group's rows together and in \
row order"
));
}
sizes.push(end - start);
at += end - start;
}
Ok(Some(sizes))
}
pub(crate) fn for_each_entry(&self, mut f: impl FnMut(usize, u32, f32)) {
for r in 0..self.n_rows {
self.for_row_entry(r, |c, v| f(r, c, v));
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct CscView {
n_rows: usize,
n_cols: usize,
col_ptr: Vec<usize>,
rows: Vec<u32>,
vals: Vec<f32>,
}
impl CscView {
#[inline]
pub fn n_rows(&self) -> usize {
self.n_rows
}
#[inline]
pub fn n_cols(&self) -> usize {
self.n_cols
}
#[inline]
pub fn column(&self, col: usize) -> (&[u32], &[f32]) {
let (s, e) = (self.col_ptr[col], self.col_ptr[col + 1]);
(&self.rows[s..e], &self.vals[s..e])
}
#[inline]
pub fn col_len(&self, col: usize) -> usize {
self.col_ptr[col + 1] - self.col_ptr[col]
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_dense() -> DMatrix {
let data = vec![1.0, 2.0, f32::NAN, 5.0, 3.0, 6.0];
DMatrix::from_dense(&data, 3, 2).unwrap()
}
#[test]
fn dense_get_and_missing() {
let d = sample_dense();
assert_eq!(d.get(0, 0), Some(1.0));
assert_eq!(d.get(1, 0), None); assert_eq!(d.get(1, 1), Some(5.0));
}
#[test]
fn row_entries_skip_missing() {
let d = sample_dense();
let mut entries = Vec::new();
d.for_row_entry(1, |index, value| entries.push((index, value)));
assert_eq!(entries, [(1, 5.0)]);
}
#[test]
fn csc_matches_dense() {
let d = sample_dense();
let csc = d.to_csc();
let (rows, vals) = csc.column(0);
assert_eq!(rows, &[0, 2]);
assert_eq!(vals, &[1.0, 3.0]);
assert_eq!(csc.col_len(1), 3);
}
#[test]
fn csr_roundtrip() {
let indptr = vec![0, 2, 3, 5];
let indices = vec![0, 1, 1, 0, 1];
let values = vec![1.0, 2.0, 5.0, 3.0, 6.0];
let d = DMatrix::from_csr(indptr, indices, values, 2).unwrap();
assert_eq!(d.get(0, 0), Some(1.0));
assert_eq!(d.get(1, 0), None);
assert_eq!(d.get(2, 1), Some(6.0));
let csc = d.to_csc();
let (rows, vals) = csc.column(0);
assert_eq!(rows, &[0, 2]);
assert_eq!(vals, &[1.0, 3.0]);
}
#[test]
fn label_length_checked() {
let d = sample_dense();
assert!(d.clone().with_labels(&[1.0, 2.0]).is_err());
assert!(d.with_labels(&[1.0, 2.0, 3.0]).is_ok());
}
#[test]
fn malformed_csr_is_rejected() {
assert!(DMatrix::from_csr(vec![1, 1], vec![], vec![], 2).is_err());
assert!(DMatrix::from_csr(vec![0, 2, 1], vec![0], vec![1.0], 2).is_err());
assert!(DMatrix::from_csr(vec![0, 2], vec![0, 0], vec![1.0, 2.0], 2).is_err());
assert!(DMatrix::from_csr(vec![0, 1], vec![0], vec![f32::INFINITY], 2).is_err());
}
#[test]
fn metadata_values_are_validated() {
let d = sample_dense();
assert!(d.clone().with_weights(&[1.0, -1.0, 1.0]).is_err());
assert!(d.clone().with_weights(&[0.0, 0.0, 0.0]).is_err());
assert!(d.clone().with_base_margin(&[0.0, 1.0]).is_err());
assert!(d.clone().with_group_sizes(&[1, 0, 2]).is_err());
assert!(d.select_rows(&[3]).is_err());
}
#[test]
fn categorical_values_must_be_non_negative_integers() {
let d = DMatrix::from_dense(&[0.0, 1.5], 2, 1).unwrap();
assert!(d.with_feature_types(&[FeatureType::Categorical]).is_err());
}
#[test]
fn label_matrix_is_row_major_and_validated() {
let d = sample_dense();
assert!(d.clone().with_label_matrix(&[1.0; 5], 2).is_err());
assert!(d.clone().with_label_matrix(&[], 0).is_err());
assert!(
d.clone()
.with_label_matrix(&[1.0, 2.0, 3.0, 4.0, 5.0, f32::NAN], 2)
.is_err()
);
let m = d
.with_label_matrix(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], 2)
.unwrap();
assert_eq!(m.n_targets(), 2);
let info = m.info();
assert_eq!((info.n_rows, info.n_targets()), (3, 2));
assert_eq!(info.label_values(), &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
assert_eq!(m.with_labels(&[0.0, 1.0, 2.0]).unwrap().n_targets(), 1);
}
#[test]
fn label_bounds_allow_censoring_but_not_nan() {
let d = sample_dense();
let inf = f32::INFINITY;
let censored = d
.clone()
.with_label_bounds(&[0.0, -1.0, 2.0], &[1.0, inf, 1.5])
.unwrap();
assert_eq!(censored.label_lower_bound().unwrap(), &[0.0, -1.0, 2.0]);
assert_eq!(censored.label_upper_bound().unwrap(), &[1.0, inf, 1.5]);
assert!(censored.labels().is_none());
assert!(censored.info().labels.is_none());
assert!(
d.clone()
.with_label_bounds(&[f32::NAN, 0.0, 0.0], &[1.0; 3])
.is_err()
);
assert!(
d.clone()
.with_label_bounds(&[0.0; 3], &[1.0, f32::NAN, 1.0])
.is_err()
);
assert!(d.with_label_bounds(&[0.0; 2], &[1.0; 3]).is_err());
}
#[test]
fn feature_weights_are_validated() {
let d = sample_dense();
assert!(d.clone().with_feature_weights(&[1.0]).is_err());
assert!(d.clone().with_feature_weights(&[1.0, -0.5]).is_err());
assert!(d.clone().with_feature_weights(&[0.0, 0.0]).is_err());
assert!(
d.clone()
.with_feature_weights(&[1.0, f32::INFINITY])
.is_err()
);
let w = d.with_feature_weights(&[0.0, 2.0]).unwrap();
assert_eq!(w.feature_weights().unwrap(), &[0.0, 2.0]);
}
#[test]
fn select_rows_carries_multi_target_labels_and_metadata() {
let d = sample_dense()
.with_label_matrix(&[10.0, 11.0, 20.0, 21.0, 30.0, 31.0], 2)
.unwrap()
.with_label_bounds(&[1.0, 2.0, 3.0], &[1.5, f32::INFINITY, 3.5])
.unwrap()
.with_feature_weights(&[0.25, 0.75])
.unwrap();
let s = d.select_rows(&[2, 0]).unwrap();
assert_eq!(s.n_targets(), 2);
assert_eq!(s.labels().unwrap(), &[30.0, 31.0, 10.0, 11.0]);
assert_eq!(s.label_lower_bound().unwrap(), &[3.0, 1.0]);
assert_eq!(s.label_upper_bound().unwrap(), &[3.5, 1.5]);
assert_eq!(s.feature_weights().unwrap(), &[0.25, 0.75]);
}
#[test]
fn select_rows_keeps_dense_storage_and_missing_entries() {
let d = DMatrix::from_dense_with_missing(&[1.0, -1.0, -1.0, 4.0, 5.0, 6.0], 3, 2, -1.0)
.unwrap();
let s = d.select_rows(&[1, 0, 1]).unwrap();
assert!(s.missing().is_nan());
let values = s.dense_values().unwrap();
assert_eq!(values.len(), 6);
let expected = [None, Some(4.0), Some(1.0), None, None, Some(4.0)];
for (cell, &want) in expected.iter().enumerate() {
assert_eq!(s.get(cell / 2, cell % 2), want, "cell {cell}");
}
assert!(d.select_rows(&[]).is_err());
}
fn is_rows_refusal(result: &Result<DMatrix>) -> bool {
matches!(
result,
Err(HessboostError::InvalidParameter { name: "rows", .. })
)
}
#[test]
fn select_rows_keeps_whole_query_groups() {
let x: Vec<f32> = (0..6).map(|i| i as f32).collect();
let d = DMatrix::from_dense(&x, 6, 1)
.unwrap()
.with_group_sizes(&[2, 3, 1])
.unwrap()
.with_group_weights(&[1.0, 2.0, 3.0])
.unwrap();
let s = d.select_rows(&[5, 0, 1, 2, 3, 4, 5]).unwrap();
assert_eq!(s.group().unwrap().group_ptr, [0, 1, 3, 6, 7]);
assert_eq!(s.weights().unwrap(), &[3.0, 1.0, 1.0, 2.0, 2.0, 2.0, 3.0]);
assert_eq!(s.get(0, 0), Some(5.0));
assert!(is_rows_refusal(&d.select_rows(&[0, 1, 2, 3])));
assert!(is_rows_refusal(&d.select_rows(&[3, 2, 4])));
assert!(is_rows_refusal(&d.select_rows(&[1, 0])));
assert!(is_rows_refusal(&d.select_rows(&[0, 5, 1])));
assert!(
sample_dense()
.select_rows(&[2, 0])
.unwrap()
.group()
.is_none()
);
}
#[test]
fn csr_categorical_validation_reports_the_first_invalid_column() {
let d = DMatrix::from_csr(
vec![0, 2, 4],
vec![1, 0, 0, 1],
vec![2.5, 1.0, -3.0, 7.5],
2,
)
.unwrap();
let types = [FeatureType::Categorical; 2];
let err = d
.clone()
.with_feature_types(&types)
.unwrap_err()
.to_string();
assert!(
err.contains("feature 0 contains invalid category value -3"),
"{err}"
);
let numeric_first = [FeatureType::Numerical, FeatureType::Categorical];
let err = d
.clone()
.with_feature_types(&numeric_first)
.unwrap_err()
.to_string();
assert!(
err.contains("feature 1 contains invalid category value 2.5"),
"{err}"
);
let d = DMatrix::from_csr(vec![0, 1, 2], vec![1, 0], vec![3.0, 1.5], 2).unwrap();
assert!(d.with_feature_types(&numeric_first).is_ok());
}
#[test]
fn group_weights_expand_across_query_rows() {
let d = sample_dense()
.with_group_sizes(&[2, 1])
.unwrap()
.with_group_weights(&[0.5, 2.0])
.unwrap();
assert_eq!(d.weights().unwrap(), &[0.5, 0.5, 2.0]);
}
}