use super::{
MultiIndex, Selection, Strategy, capped, index_fn_dim, index_fn_supports_lookup,
multi_index_to_flat,
};
use crate::iteration::comprehension::metadata::{IndexFn, cycle_length};
use crate::iteration::comprehension::strategy::StrategyName;
pub struct Extrema;
impl Strategy for Extrema {
fn name(&self) -> StrategyName {
StrategyName::Extrema
}
fn selects_from_shape(&self) -> bool {
true
}
fn accepts_input(&self, idx: Option<&IndexFn>) -> bool {
idx.is_some()
}
fn has_closed_form_for(&self, idx: &IndexFn) -> bool {
matches!(
idx,
IndexFn::Lattice { .. } | IndexFn::Continuous { .. } | IndexFn::Hybrid { .. }
)
}
fn select(
&self,
index_fn: &IndexFn,
cardinality: u64,
truncation: Option<u64>,
_seed: Option<u64>,
) -> Selection {
if index_fn_supports_lookup(index_fn) {
let mis = extrema_multi_indices(index_fn, truncation);
Selection::from_multi_indices(index_fn, mis, cardinality)
} else {
Selection::Positions(naive_extrema_positions(cardinality, truncation))
}
}
fn select_surviving(
&self,
index_fn: &IndexFn,
cardinality: u64,
truncation: Option<u64>,
seed: Option<u64>,
survivors: &[u64],
) -> Selection {
if !index_fn_supports_lookup(index_fn) || index_fn_dim(index_fn) == 0 {
return super::surviving_in_rank(
&|count| self.select(index_fn, cardinality, count, seed),
cardinality,
truncation,
survivors,
);
}
let kept: Vec<(u64, MultiIndex)> = extrema_scored(&extrema_axis_sizes(index_fn))
.into_iter()
.filter(|(_, mi)| {
multi_index_to_flat(index_fn, mi)
.is_some_and(|p| survivors.binary_search(&(p as u64)).is_ok())
})
.collect();
Selection::from_multi_indices(index_fn, take_n_strata(kept, truncation), cardinality)
}
}
fn naive_extrema_positions(total: u64, truncation: Option<u64>) -> Vec<u64> {
let n = capped(truncation, total);
if n == 0 {
return Vec::new();
}
if n == 1 {
return vec![0];
}
let mut out = Vec::with_capacity(n as usize);
out.push(0);
out.push(total - 1);
out.extend((1..total - 1).take((n - 2) as usize));
out
}
pub(crate) fn extrema_multi_indices(idx: &IndexFn, truncation: Option<u64>) -> Vec<MultiIndex> {
if index_fn_dim(idx) == 0 {
return Vec::new();
}
extrema_strata(&extrema_axis_sizes(idx), truncation)
}
fn extrema_axis_sizes(idx: &IndexFn) -> Vec<u64> {
match idx {
IndexFn::Lattice { axis_sizes } => axis_sizes.clone(),
IndexFn::Continuous { intervals, .. } => vec![2u64; intervals.len()],
IndexFn::Hybrid {
discrete_axes,
continuous_axes,
..
} => {
let mut s = Vec::with_capacity(discrete_axes.len() + continuous_axes.len());
s.extend(discrete_axes.iter().copied());
s.extend(continuous_axes.iter().map(|_| 2u64));
s
}
IndexFn::Lockstep { length } => vec![*length],
IndexFn::Modular { axis_sizes } => vec![cycle_length(axis_sizes)],
IndexFn::Concatenation { segment_sizes } => {
vec![segment_sizes.iter().copied().sum()]
}
}
}
fn interior_count(mi: &[u64], axis_sizes: &[u64]) -> u64 {
let mut count = 0;
for (c, s) in mi.iter().zip(axis_sizes.iter()) {
if *c != 0 && *c != s.saturating_sub(1) {
count += 1;
}
}
count
}
fn extrema_strata(axis_sizes: &[u64], truncation: Option<u64>) -> Vec<MultiIndex> {
take_n_strata(extrema_scored(axis_sizes), truncation)
}
fn extrema_scored(axis_sizes: &[u64]) -> Vec<(u64, MultiIndex)> {
let total: u64 = axis_sizes.iter().product();
if total == 0 {
return Vec::new();
}
let mut scored: Vec<(u64, MultiIndex)> = Vec::with_capacity(total as usize);
let mut mi = vec![0u64; axis_sizes.len()];
for _ in 0..total {
scored.push((interior_count(&mi, axis_sizes), mi.clone()));
for axis in (0..axis_sizes.len()).rev() {
mi[axis] += 1;
if mi[axis] < axis_sizes[axis] {
break;
}
mi[axis] = 0;
}
}
scored.sort_by(|(ia, a), (ib, b)| ia.cmp(ib).then_with(|| a.cmp(b)));
scored
}
fn take_n_strata(scored: Vec<(u64, MultiIndex)>, n: Option<u64>) -> Vec<MultiIndex> {
let limit = match n {
None => return scored.into_iter().map(|(_, mi)| mi).collect(),
Some(0) => return Vec::new(),
Some(k) => k,
};
let mut out = Vec::with_capacity(scored.len());
let mut strata_seen = 0u64;
let mut last: Option<u64> = None;
for (key, mi) in scored {
if last != Some(key) {
if strata_seen >= limit {
break;
}
strata_seen += 1;
last = Some(key);
}
out.push(mi);
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extrema_2x2_emits_4_corners() {
let idx = IndexFn::Lattice {
axis_sizes: vec![2, 2],
};
let out = extrema_multi_indices(&idx, None);
assert_eq!(out.len(), 4);
let mut sorted = out.clone();
sorted.sort();
assert_eq!(sorted, vec![vec![0, 0], vec![0, 1], vec![1, 0], vec![1, 1]]);
}
#[test]
fn a_3x3x3_lattice_strata_are_corners_edges_faces_then_interior() {
let idx = IndexFn::Lattice {
axis_sizes: vec![3, 3, 3],
};
assert_eq!(extrema_multi_indices(&idx, Some(1)).len(), 8); assert_eq!(extrema_multi_indices(&idx, Some(2)).len(), 20); assert_eq!(extrema_multi_indices(&idx, Some(3)).len(), 26); assert_eq!(extrema_multi_indices(&idx, Some(4)).len(), 27); assert_eq!(extrema_multi_indices(&idx, None).len(), 27);
assert_eq!(
extrema_multi_indices(&idx, Some(1)),
vec![
vec![0, 0, 0],
vec![0, 0, 2],
vec![0, 2, 0],
vec![0, 2, 2],
vec![2, 0, 0],
vec![2, 0, 2],
vec![2, 2, 0],
vec![2, 2, 2],
]
);
assert_eq!(
extrema_multi_indices(&idx, None).last(),
Some(&vec![1u64, 1, 1])
);
}
#[test]
fn extrema_3x3_edges_then_center() {
let idx = IndexFn::Lattice {
axis_sizes: vec![3, 3],
};
assert_eq!(extrema_multi_indices(&idx, Some(1)).len(), 4);
assert_eq!(extrema_multi_indices(&idx, Some(2)).len(), 8);
assert_eq!(extrema_multi_indices(&idx, None).len(), 9);
let corners = extrema_multi_indices(&idx, Some(1));
let mut sorted = corners.clone();
sorted.sort();
assert_eq!(sorted, vec![vec![0, 0], vec![0, 2], vec![2, 0], vec![2, 2]]);
}
#[test]
fn extrema_partial_count_keeps_whole_stratum() {
let idx = IndexFn::Lattice {
axis_sizes: vec![2, 2],
};
for n in [1u64, 2, 3, 4] {
let out = extrema_multi_indices(&idx, Some(n));
assert_eq!(out.len(), 4, "extrema/{n} should keep the whole stratum");
}
let idx3 = IndexFn::Lattice {
axis_sizes: vec![2, 2, 2],
};
assert_eq!(extrema_multi_indices(&idx3, Some(1)).len(), 8);
}
#[test]
fn extrema_zero_strata_is_empty() {
let idx = IndexFn::Lattice {
axis_sizes: vec![2, 2],
};
assert!(extrema_multi_indices(&idx, Some(0)).is_empty());
}
#[test]
fn extrema_1d_partial_keeps_both_endpoints() {
let idx = IndexFn::Lattice {
axis_sizes: vec![5],
};
let out = extrema_multi_indices(&idx, Some(1));
let mut sorted = out.clone();
sorted.sort();
assert_eq!(sorted, vec![vec![0], vec![4]]);
}
#[test]
fn extrema_1d_two_strata() {
let idx = IndexFn::Lattice {
axis_sizes: vec![5],
};
let mut endpoints = extrema_multi_indices(&idx, Some(1));
endpoints.sort();
assert_eq!(endpoints, vec![vec![0], vec![4]]);
assert_eq!(extrema_multi_indices(&idx, None).len(), 5);
assert_eq!(extrema_multi_indices(&idx, None)[0..2], [vec![0], vec![4]]);
}
#[test]
fn extrema_3d_8_corners() {
let idx = IndexFn::Lattice {
axis_sizes: vec![2, 2, 2],
};
let out = extrema_multi_indices(&idx, None);
assert_eq!(out.len(), 8);
}
#[test]
fn continuous_box_corners() {
use crate::iteration::comprehension::cardinality::{Interval, ProductMeasure};
let idx = IndexFn::Continuous {
intervals: vec![Interval::closed(0.0, 1.0), Interval::closed(-1.0, 1.0)],
measure: ProductMeasure::Uniform,
};
let out = extrema_multi_indices(&idx, None);
assert_eq!(out.len(), 4);
let mut sorted = out.clone();
sorted.sort();
assert_eq!(sorted, vec![vec![0, 0], vec![0, 1], vec![1, 0], vec![1, 1]]);
}
#[test]
fn lockstep_endpoints_then_interior() {
let idx = IndexFn::Lockstep { length: 10 };
assert_eq!(extrema_multi_indices(&idx, Some(1)), vec![vec![0], vec![9]]);
let all = extrema_multi_indices(&idx, None);
assert_eq!(all.len(), 10);
assert_eq!(all[0..2], [vec![0], vec![9]]);
}
}