use crate::matrix::rand_util::mix_seed;
use rand::prelude::SliceRandom;
use rand::rngs::SmallRng;
use rand::SeedableRng;
use rayon::prelude::*;
use rustc_hash::FxHashMap as HashMap;
use std::hash::Hash;
const PARTITION_SHUFFLE_SEED: u64 = 0x5041_5254_5348_5546;
pub fn median(values: &[f32]) -> f32 {
if values.is_empty() {
return 0.0;
}
let mut sorted = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = sorted.len();
if n % 2 == 1 {
sorted[n / 2]
} else {
0.5 * (sorted[n / 2 - 1] + sorted[n / 2])
}
}
pub fn quantiles(values: &[f32], qs: &[f64]) -> Vec<f32> {
if values.is_empty() {
return vec![0.0; qs.len()];
}
let mut sorted = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let last = sorted.len() - 1;
qs.iter()
.map(|&q| sorted[((q.clamp(0.0, 1.0) * last as f64).floor() as usize).min(last)])
.collect()
}
pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na > 0.0 && nb > 0.0 {
dot / (na * nb)
} else {
0.0
}
}
pub fn partition_by_membership<T>(
membership: &[T],
nelem_per_group: Option<usize>,
) -> HashMap<T, Vec<usize>>
where
T: Eq + Hash + Clone + Send + Sync,
{
let mut pb_elems: HashMap<T, Vec<usize>> = HashMap::default();
for (cell, k) in membership.iter().enumerate() {
pb_elems.entry(k.clone()).or_default().push(cell);
}
pb_elems.par_iter_mut().for_each(|(_k, cells)| {
let ncells = cells.len();
if let Some(ntarget) = nelem_per_group {
if ncells > ntarget {
let key = cells.first().copied().unwrap_or(0) as u64;
let mut rng = SmallRng::seed_from_u64(mix_seed(PARTITION_SHUFFLE_SEED, key));
cells.shuffle(&mut rng);
cells.truncate(ntarget);
}
}
});
pb_elems
}
pub fn default_block_size(num_features: usize) -> usize {
const MIN_BLOCK_SIZE: usize = 100;
const MAX_BLOCK_SIZE: usize = 10_000;
const TARGET_WORK: usize = MIN_BLOCK_SIZE * 10_000;
if num_features == 0 {
return MIN_BLOCK_SIZE;
}
(TARGET_WORK / num_features).clamp(MIN_BLOCK_SIZE, MAX_BLOCK_SIZE)
}
pub fn generate_minibatch_intervals(
ntot: usize,
num_features: usize,
batch_size: Option<usize>,
) -> Vec<(usize, usize)> {
let batch_size = batch_size.unwrap_or_else(|| default_block_size(num_features));
let num_batches = ntot.div_ceil(batch_size);
(0..num_batches)
.map(|b| {
let lb: usize = b * batch_size;
let ub: usize = ((b + 1) * batch_size).min(ntot);
(lb, ub)
})
.collect::<Vec<_>>()
}
pub fn byte_budget_intervals(
nnz_per_col: &[u64],
budget_bytes: usize,
bytes_per_nnz: usize,
) -> Vec<(usize, usize)> {
let budget_nnz = (budget_bytes / bytes_per_nnz.max(1)).max(1) as u64;
let mut intervals = Vec::new();
let mut lb = 0usize;
let mut acc = 0u64;
for (col, &nnz) in nnz_per_col.iter().enumerate() {
if acc > 0 && acc + nnz > budget_nnz {
intervals.push((lb, col));
lb = col;
acc = 0;
}
acc += nnz;
}
if lb < nnz_per_col.len() {
intervals.push((lb, nnz_per_col.len()));
}
intervals
}
#[cfg(test)]
#[path = "utils_tests.rs"]
mod tests;