use std::cmp::Ordering;
use crate::core::number::Number;
use super::knn_sieve;
pub fn find_kth_d<'a, T: Number, U: Number>(
grains: &mut [knn_sieve::Grain<'a, T, U>],
cumulative_cardinalities: &mut [usize],
k: usize,
) -> (knn_sieve::Grain<'a, T, U>, usize) {
assert_eq!(grains.len(), cumulative_cardinalities.len());
_find_kth(
grains,
cumulative_cardinalities,
k,
0,
grains.len() - 1,
&knn_sieve::Delta::D,
)
}
pub fn find_kth_d_max<'a, T: Number, U: Number>(
grains: &mut [knn_sieve::Grain<'a, T, U>],
cumulative_cardinalities: &mut [usize],
k: usize,
) -> (knn_sieve::Grain<'a, T, U>, usize) {
assert_eq!(grains.len(), cumulative_cardinalities.len());
_find_kth(
grains,
cumulative_cardinalities,
k,
0,
grains.len() - 1,
&knn_sieve::Delta::Max,
)
}
pub fn find_kth_d_min<'a, T: Number, U: Number>(
grains: &mut [knn_sieve::Grain<'a, T, U>],
cumulative_cardinalities: &mut [usize],
k: usize,
) -> (knn_sieve::Grain<'a, T, U>, usize) {
assert_eq!(grains.len(), cumulative_cardinalities.len());
_find_kth(
grains,
cumulative_cardinalities,
k,
0,
grains.len() - 1,
&knn_sieve::Delta::Min,
)
}
pub fn _find_kth<'a, T: Number, U: Number>(
grains: &mut [knn_sieve::Grain<'a, T, U>],
cumulative_cardinalities: &mut [usize],
k: usize,
l: usize,
r: usize,
delta: &knn_sieve::Delta,
) -> (knn_sieve::Grain<'a, T, U>, usize) {
let mut cardinalities = (0..1)
.chain(cumulative_cardinalities.iter().copied())
.zip(cumulative_cardinalities.iter().copied())
.map(|(prev, next)| next - prev)
.collect::<Vec<_>>();
assert_eq!(cardinalities.len(), grains.len());
let position = partition(grains, &mut cardinalities, l, r, delta);
let mut cumulative_cardinalities = cardinalities
.iter()
.scan(0, |acc, v| {
*acc += *v;
Some(*acc)
})
.collect::<Vec<_>>();
match cumulative_cardinalities[position].cmp(&k) {
Ordering::Less => _find_kth(grains, &mut cumulative_cardinalities, k, position + 1, r, delta),
Ordering::Equal => (grains[position].clone(), position),
Ordering::Greater => {
if (position > 0) && (cumulative_cardinalities[position - 1] > k) {
_find_kth(grains, &mut cumulative_cardinalities, k, l, position - 1, delta)
} else {
(grains[position].clone(), position)
}
}
}
}
pub fn find_kth_threshold<U: Number>(thresholds: &mut [(U, usize)], k: usize) -> (usize, (U, usize)) {
let (index, (threshold, _)) = _find_kth_threshold(thresholds, k, 0, thresholds.len() - 1);
let mut b = index;
for a in (index + 1)..(thresholds.len()) {
if thresholds[a].0 == threshold {
b += 1;
thresholds.swap(a, b);
}
}
(b, thresholds[b])
}
fn _find_kth_threshold<U: Number>(thresholds: &mut [(U, usize)], k: usize, l: usize, r: usize) -> (usize, (U, usize)) {
if l >= r {
let position = std::cmp::min(l, r);
(position, thresholds[position])
} else {
let position = partition_threshold(thresholds, l, r);
let guaranteed_cardinalities = thresholds
.iter()
.scan(0, |acc, &(_, v)| {
*acc += v;
Some(*acc)
})
.collect::<Vec<_>>();
assert!(
guaranteed_cardinalities[r] > k,
"Too few guarantees {} vs {} at {} ...",
guaranteed_cardinalities[r],
k,
r
);
let num_guaranteed = guaranteed_cardinalities[position];
match num_guaranteed.cmp(&k) {
Ordering::Less => _find_kth_threshold(thresholds, k, position + 1, r),
Ordering::Equal => (position, thresholds[position]),
Ordering::Greater => {
if (position > 0) && (guaranteed_cardinalities[position - 1] > k) {
_find_kth_threshold(thresholds, k, l, position - 1)
} else {
(position, thresholds[position])
}
}
}
}
}
pub fn partition_threshold<U: Number>(thresholds: &mut [(U, usize)], l: usize, r: usize) -> usize {
let pivot = (l + r) / 2; thresholds.swap(pivot, r);
let (mut a, mut b) = (l, l);
while b < r {
if thresholds[b].0 <= thresholds[r].0 {
thresholds.swap(a, b);
a += 1;
}
b += 1;
}
thresholds.swap(a, r);
a
}
fn partition<T: Number, U: Number>(
grains: &mut [knn_sieve::Grain<T, U>],
cardinalities: &mut [usize],
l: usize,
r: usize,
delta: &knn_sieve::Delta,
) -> usize {
let pivot = (l + r) / 2; swaps(grains, cardinalities, pivot, r);
let (mut a, mut b) = (l, l);
while b < r {
let grain_comp = match delta {
knn_sieve::Delta::D => (grains[b]).ord_by_d(&grains[r]),
knn_sieve::Delta::Max => (grains[b]).ord_by_d_max(&grains[r]),
knn_sieve::Delta::Min => (grains[b]).ord_by_d_min(&grains[r]),
};
if grain_comp == Ordering::Less {
swaps(grains, cardinalities, a, b);
a += 1;
}
b += 1;
}
swaps(grains, cardinalities, a, r);
a
}
pub fn swaps<T>(grains: &mut [T], cardinalities: &mut [usize], i: usize, j: usize) {
grains.swap(i, j);
cardinalities.swap(i, j);
}