extern crate alloc;
use alloc::vec::Vec;
use crate::constants::DECILES;
use crate::math;
use crate::types::Vector9;
#[inline]
fn decile_indices(n: usize, k: usize) -> (usize, usize) {
debug_assert!((1..=9).contains(&k), "decile k must be in 1..=9");
debug_assert!(n > 0, "n must be positive");
debug_assert!(
n <= (usize::MAX - 5) / 9,
"n too large, would overflow in index computation"
);
let h_numerator = n * k + 5;
let floor_h = h_numerator / 10; let has_fraction = !h_numerator.is_multiple_of(10);
let ceil_h = if has_fraction { floor_h + 1 } else { floor_h };
let floor_idx = floor_h.saturating_sub(1).min(n - 1);
let ceil_idx = ceil_h.saturating_sub(1).min(n - 1);
(floor_idx, ceil_idx)
}
#[inline]
fn debug_assert_finite(data: &[f64]) {
debug_assert!(
data.iter().all(|x| x.is_finite()),
"quantile input must be finite (no NaN or infinity)"
);
}
pub fn compute_quantile(data: &mut [f64], p: f64) -> f64 {
assert!(!data.is_empty(), "Cannot compute quantile of empty slice");
assert!(
(0.0..=1.0).contains(&p),
"Quantile probability must be in [0, 1]"
);
debug_assert_finite(data);
let n = data.len();
if n == 1 {
return data[0];
}
let h = n as f64 * p + 0.5;
let floor_idx = (math::floor(h) as usize).saturating_sub(1).min(n - 1);
let ceil_idx = (math::ceil(h) as usize).saturating_sub(1).min(n - 1);
let cmp = |a: &f64, b: &f64| a.total_cmp(b);
if floor_idx == ceil_idx {
let (_, mid, _) = data.select_nth_unstable_by(floor_idx, cmp);
return *mid;
}
let (_, mid, _) = data.select_nth_unstable_by(ceil_idx, cmp);
let ceil_val = *mid;
let (_, mid, _) = data[..ceil_idx].select_nth_unstable_by(floor_idx, cmp);
let floor_val = *mid;
(floor_val + ceil_val) / 2.0
}
pub fn compute_deciles(data: &[f64]) -> Vector9 {
assert!(!data.is_empty(), "Cannot compute deciles of empty slice");
debug_assert_finite(data);
let mut sorted = data.to_vec();
sorted.sort_unstable_by(|a, b| a.total_cmp(b));
compute_deciles_sorted(&sorted)
}
pub fn compute_deciles_fast(data: &[f64]) -> Vector9 {
assert!(!data.is_empty(), "Cannot compute deciles of empty slice");
debug_assert_finite(data);
let mut working = data.to_vec();
working.sort_unstable_by(|a, b| a.total_cmp(b));
compute_deciles_sorted(&working)
}
pub fn compute_deciles_with_buffer(data: &[f64], buffer: &mut Vec<f64>) -> Vector9 {
assert!(!data.is_empty(), "Cannot compute deciles of empty slice");
debug_assert_finite(data);
buffer.clear();
buffer.extend_from_slice(data);
buffer.sort_unstable_by(|a, b| a.total_cmp(b));
compute_deciles_sorted(buffer)
}
pub fn compute_deciles_inplace(data: &mut [f64]) -> Vector9 {
assert!(!data.is_empty(), "Cannot compute deciles of empty slice");
debug_assert_finite(data);
let n = data.len();
if n <= 200 {
data.sort_unstable_by(|a, b| a.total_cmp(b));
return compute_deciles_sorted(data);
}
let mut decile_idx_pairs = [(0usize, 0usize); 9];
for k in 1..=9 {
decile_idx_pairs[k - 1] = decile_indices(n, k);
}
let mut indices = [0usize; 18];
let mut num_indices = 0;
for &(floor_idx, ceil_idx) in &decile_idx_pairs {
if num_indices == 0 || indices[num_indices - 1] != floor_idx {
indices[num_indices] = floor_idx;
num_indices += 1;
}
if ceil_idx != floor_idx && (num_indices == 0 || indices[num_indices - 1] != ceil_idx) {
indices[num_indices] = ceil_idx;
num_indices += 1;
}
}
debug_assert!(
indices[..num_indices].windows(2).all(|w| w[0] < w[1]),
"decile indices must be strictly increasing"
);
multi_select(data, &indices[..num_indices]);
let mut result = Vector9::zeros();
for (i, &(floor_idx, ceil_idx)) in decile_idx_pairs.iter().enumerate() {
result[i] = (data[floor_idx] + data[ceil_idx]) / 2.0;
}
result
}
fn multi_select(data: &mut [f64], indices: &[usize]) {
if indices.is_empty() {
return;
}
multi_select_recursive(data, indices, 0, data.len());
}
fn multi_select_recursive(data: &mut [f64], indices: &[usize], lo: usize, hi: usize) {
if indices.is_empty() || hi.saturating_sub(lo) <= 1 {
return;
}
debug_assert!(indices[0] >= lo);
debug_assert!(indices[indices.len() - 1] < hi);
if indices.len() == 1 {
let target = indices[0];
let rel_idx = target - lo;
data[lo..hi].select_nth_unstable_by(rel_idx, |a, b| a.total_cmp(b));
return;
}
let mid = indices.len() / 2;
let pivot_abs = indices[mid];
let pivot_rel = pivot_abs - lo;
data[lo..hi].select_nth_unstable_by(pivot_rel, |a, b| a.total_cmp(b));
if mid > 0 {
multi_select_recursive(data, &indices[..mid], lo, pivot_abs);
}
if mid + 1 < indices.len() {
multi_select_recursive(data, &indices[mid + 1..], pivot_abs + 1, hi);
}
}
pub fn compute_deciles_sorted(sorted: &[f64]) -> Vector9 {
assert!(!sorted.is_empty(), "Cannot compute deciles of empty slice");
let n = sorted.len();
let mut result = Vector9::zeros();
for k in 1..=9 {
let (floor_idx, ceil_idx) = decile_indices(n, k);
result[k - 1] = (sorted[floor_idx] + sorted[ceil_idx]) / 2.0;
}
result
}
pub fn compute_midquantile(data: &mut [f64], p: f64) -> f64 {
assert!(!data.is_empty(), "Cannot compute quantile of empty slice");
assert!(
(0.0..=1.0).contains(&p),
"Quantile probability must be in [0, 1]"
);
debug_assert_finite(data);
let n = data.len();
data.sort_by(|a, b| a.total_cmp(b));
if n == 1 {
return data[0];
}
let mut i = 0;
while i < n {
let value = data[i];
let mut count = 1;
while i + count < n && data[i + count] == value {
count += 1;
}
let f_mid = (i as f64 + count as f64 / 2.0) / n as f64;
if p <= f_mid {
return value;
}
i += count;
}
data[n - 1]
}
pub fn compute_midquantile_deciles(data: &[f64]) -> Vector9 {
assert!(!data.is_empty(), "Cannot compute deciles of empty slice");
debug_assert_finite(data);
let mut sorted = data.to_vec();
sorted.sort_by(|a, b| a.total_cmp(b));
compute_midquantile_deciles_sorted(&sorted)
}
pub fn compute_midquantile_deciles_sorted(sorted: &[f64]) -> Vector9 {
assert!(!sorted.is_empty(), "Cannot compute deciles of empty slice");
let n = sorted.len();
let mut result = Vector9::zeros();
for (i, &p) in DECILES.iter().enumerate() {
result[i] = midquantile_from_sorted(sorted, n, p);
}
result
}
fn midquantile_from_sorted(sorted: &[f64], n: usize, p: f64) -> f64 {
if n == 1 {
return sorted[0];
}
let mut i = 0;
while i < n {
let value = sorted[i];
let mut count = 1;
while i + count < n && sorted[i + count] == value {
count += 1;
}
let f_mid = (i as f64 + count as f64 / 2.0) / n as f64;
if p <= f_mid {
return value;
}
i += count;
}
sorted[n - 1]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_compute_quantile_median() {
let mut data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let median = compute_quantile(&mut data, 0.5);
assert!((median - 3.0).abs() < 1e-10);
}
#[test]
fn test_compute_quantile_extremes() {
let mut data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let min = compute_quantile(&mut data.clone(), 0.0);
let max = compute_quantile(&mut data, 1.0);
assert!((min - 1.0).abs() < 1e-10, "min was {}", min);
assert!((max - 5.0).abs() < 1e-10, "max was {}", max);
}
#[test]
fn test_compute_deciles_sorted() {
let data: Vec<f64> = (1..=100).map(|x| x as f64).collect();
let deciles = compute_deciles(&data);
for i in 1..9 {
assert!(deciles[i] >= deciles[i - 1]);
}
}
#[test]
fn test_compute_deciles_fast_matches_sort() {
let data: Vec<f64> = (1..=100).map(|x| x as f64).collect();
let deciles_sort = compute_deciles(&data);
let deciles_fast = compute_deciles_fast(&data);
for i in 0..9 {
let diff = (deciles_sort[i] - deciles_fast[i]).abs();
assert!(
diff < 1e-10,
"Decile {} differs: sort={}, fast={}, diff={}",
i,
deciles_sort[i],
deciles_fast[i],
diff
);
}
}
#[test]
fn test_compute_deciles_fast_random_data() {
let data: Vec<f64> = vec![
3.7, 1.2, 9.5, 2.1, 7.3, 4.8, 6.2, 8.9, 1.5, 5.4, 2.7, 9.1, 3.3, 6.8, 4.5, 7.9, 2.4,
8.3, 5.7, 1.9,
];
let deciles_sort = compute_deciles(&data);
let deciles_fast = compute_deciles_fast(&data);
for i in 0..9 {
let diff = (deciles_sort[i] - deciles_fast[i]).abs();
assert!(
diff < 1e-10,
"Decile {} differs: sort={}, fast={}, diff={}",
i,
deciles_sort[i],
deciles_fast[i],
diff
);
}
for i in 1..9 {
assert!(deciles_fast[i] >= deciles_fast[i - 1]);
}
}
#[test]
fn test_compute_deciles_fast_small_data() {
let data: Vec<f64> = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let deciles_sort = compute_deciles(&data);
let deciles_fast = compute_deciles_fast(&data);
for i in 0..9 {
let diff = (deciles_sort[i] - deciles_fast[i]).abs();
assert!(
diff < 1e-10,
"Decile {} differs: sort={}, fast={}, diff={}",
i,
deciles_sort[i],
deciles_fast[i],
diff
);
}
}
#[test]
fn test_compute_deciles_fast_large_data() {
let data: Vec<f64> = (0..20000).map(|x| (x as f64 * 1.234) % 1000.0).collect();
let deciles_sort = compute_deciles(&data);
let deciles_fast = compute_deciles_fast(&data);
for i in 0..9 {
let diff = (deciles_sort[i] - deciles_fast[i]).abs();
assert!(
diff < 1e-8, "Decile {} differs: sort={}, fast={}, diff={}",
i,
deciles_sort[i],
deciles_fast[i],
diff
);
}
}
#[test]
#[should_panic(expected = "Cannot compute quantile of empty slice")]
fn test_empty_slice_panics() {
let mut data: Vec<f64> = vec![];
compute_quantile(&mut data, 0.5);
}
#[test]
#[should_panic(expected = "Cannot compute deciles of empty slice")]
fn test_compute_deciles_fast_empty_panics() {
let data: Vec<f64> = vec![];
compute_deciles_fast(&data);
}
#[test]
fn test_midquantile_no_ties() {
let mut data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let median = compute_midquantile(&mut data, 0.5);
assert!((median - 6.0).abs() < 1e-10, "Median was {}", median);
}
#[test]
fn test_midquantile_with_ties() {
let mut data = vec![1.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 3.0, 4.0];
let median = compute_midquantile(&mut data, 0.5);
assert!((median - 3.0).abs() < 1e-10, "Median was {}", median);
}
#[test]
fn test_midquantile_all_same() {
let mut data = vec![42.0; 100];
let median = compute_midquantile(&mut data, 0.5);
assert!((median - 42.0).abs() < 1e-10);
let q10 = compute_midquantile(&mut data, 0.1);
assert!((q10 - 42.0).abs() < 1e-10);
let q90 = compute_midquantile(&mut data, 0.9);
assert!((q90 - 42.0).abs() < 1e-10);
}
#[test]
fn test_midquantile_deciles_discrete_data() {
let data: Vec<f64> = (0..1000)
.map(|i| ((i % 5) * 10) as f64) .collect();
let deciles = compute_midquantile_deciles(&data);
for i in 1..9 {
assert!(
deciles[i] >= deciles[i - 1],
"Deciles should be monotonic: d[{}]={} < d[{}]={}",
i - 1,
deciles[i - 1],
i,
deciles[i]
);
}
}
#[test]
fn test_midquantile_single_element() {
let mut data = vec![42.0];
let result = compute_midquantile(&mut data, 0.5);
assert!((result - 42.0).abs() < 1e-10);
}
#[test]
fn test_midquantile_deciles_matches_sorted() {
let data: Vec<f64> = vec![1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 4.0, 4.0, 5.0, 5.0];
let deciles1 = compute_midquantile_deciles(&data);
let mut sorted = data.clone();
sorted.sort_by(|a, b| a.total_cmp(b));
let deciles2 = compute_midquantile_deciles_sorted(&sorted);
for i in 0..9 {
assert!(
(deciles1[i] - deciles2[i]).abs() < 1e-10,
"Decile {} mismatch: {} vs {}",
i,
deciles1[i],
deciles2[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_sequential() {
let data: Vec<f64> = (1..=1000).map(|x| x as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_random() {
let data: Vec<f64> = (0..2500)
.map(|i| (i as f64 * 17.3 + 42.7) % 1000.0)
.collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_reversed() {
let data: Vec<f64> = (0..1000).rev().map(|x| x as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_all_equal() {
let data: Vec<f64> = vec![42.0; 500];
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_heavy_ties() {
let data: Vec<f64> = (0..1000).map(|i| (i % 5) as f64 * 10.0).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
#[test]
fn test_multiselect_vs_sort_small_sizes() {
for n in [51, 52, 100, 101, 200] {
let data: Vec<f64> = (0..n).map(|i| i as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"n={}, Decile {} mismatch: sort={}, multiselect={}",
n,
i,
reference[i],
result[i]
);
}
}
}
#[test]
fn test_multiselect_vs_sort_bootstrap_size() {
let data: Vec<f64> = (0..5000)
.map(|i| 100.0 + ((i as f64 * std::f64::consts::PI).sin() * 50.0))
.collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-8,
"Decile {} mismatch: sort={}, multiselect={}",
i,
reference[i],
result[i]
);
}
}
}
#[cfg(test)]
mod proptests {
use super::*;
use proptest::prelude::*;
fn data_strategy(min_size: usize, max_size: usize) -> impl Strategy<Value = Vec<f64>> {
prop::collection::vec(prop::num::f64::NORMAL, min_size..=max_size)
}
fn discrete_data_strategy(min_size: usize, max_size: usize) -> impl Strategy<Value = Vec<f64>> {
prop::collection::vec(0i32..100, min_size..=max_size)
.prop_map(|v| v.into_iter().map(|x| x as f64).collect())
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(200))]
#[test]
fn prop_multiselect_matches_sort(data in data_strategy(51, 10000)) {
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
let diff = (reference[i] - result[i]).abs();
prop_assert!(
diff < 1e-10,
"Decile {} mismatch: sort={}, multiselect={}, diff={}",
i, reference[i], result[i], diff
);
}
}
#[test]
fn prop_multiselect_discrete_data(data in discrete_data_strategy(51, 5000)) {
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
let diff = (reference[i] - result[i]).abs();
prop_assert!(
diff < 1e-10,
"Decile {} mismatch with discrete data: sort={}, multiselect={}",
i, reference[i], result[i]
);
}
}
#[test]
fn prop_multiselect_small_arrays(data in data_strategy(1, 50)) {
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
let r = reference[i];
let t = result[i];
let equal = r.total_cmp(&t).is_eq() || (r - t).abs() < 1e-10;
prop_assert!(
equal,
"Decile {} mismatch on small array (n={}): sort={}, multiselect={}",
i, data.len(), r, t
);
}
}
#[test]
fn prop_multiselect_monotonic(data in data_strategy(51, 5000)) {
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 1..9 {
prop_assert!(
result[i] >= result[i - 1],
"Deciles not monotonic: d[{}]={} < d[{}]={}",
i - 1, result[i - 1], i, result[i]
);
}
}
#[test]
fn prop_multiselect_within_range(data in data_strategy(51, 5000)) {
let min_val = data.iter().cloned().fold(f64::INFINITY, f64::min);
let max_val = data.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
prop_assert!(
result[i] >= min_val && result[i] <= max_val,
"Decile {} ({}) outside data range [{}, {}]",
i, result[i], min_val, max_val
);
}
}
#[test]
fn prop_deciles_fast_matches_stable(data in data_strategy(1, 5000)) {
let stable = compute_deciles(&data);
let fast = compute_deciles_fast(&data);
for i in 0..9 {
let diff = (stable[i] - fast[i]).abs();
prop_assert!(
diff < 1e-10,
"Decile {} mismatch between stable and fast: {} vs {}",
i, stable[i], fast[i]
);
}
}
#[test]
fn prop_deciles_buffer_matches(data in data_strategy(1, 5000)) {
let reference = compute_deciles(&data);
let mut buffer = Vec::new();
let result = compute_deciles_with_buffer(&data, &mut buffer);
for i in 0..9 {
let diff = (reference[i] - result[i]).abs();
prop_assert!(
diff < 1e-10,
"Decile {} mismatch with buffer version: {} vs {}",
i, reference[i], result[i]
);
}
}
}
#[test]
fn test_multiselect_sorted_ascending() {
let data: Vec<f64> = (0..1000).map(|x| x as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Sorted ascending: decile {} mismatch",
i
);
}
}
#[test]
fn test_multiselect_sorted_descending() {
let data: Vec<f64> = (0..1000).rev().map(|x| x as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Sorted descending: decile {} mismatch",
i
);
}
}
#[test]
fn test_multiselect_all_same_value() {
let data: Vec<f64> = vec![42.0; 1000];
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"All same value: decile {} mismatch",
i
);
}
}
#[test]
fn test_multiselect_two_values() {
let data: Vec<f64> = (0..1000)
.map(|i| if i % 2 == 0 { 0.0 } else { 1.0 })
.collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Two values: decile {} mismatch",
i
);
}
}
#[test]
fn test_multiselect_boundary_size_51() {
let data: Vec<f64> = (0..51).map(|x| x as f64).collect();
let reference = compute_deciles(&data);
let mut test_data = data.clone();
let result = compute_deciles_inplace(&mut test_data);
for i in 0..9 {
assert!(
(reference[i] - result[i]).abs() < 1e-10,
"Boundary size 51: decile {} mismatch",
i
);
}
}
}