use super::{
EvaluatedInput, MultiIndex, Strategy, Tuple, index_fn_dim,
index_fn_supports_lookup, multi_index_to_flat,
};
use crate::iteration::comprehension::metadata::IndexFn;
use crate::iteration::comprehension::strategy::StrategyName;
pub struct Extrema;
impl Strategy for Extrema {
fn name(&self) -> StrategyName {
StrategyName::Extrema
}
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 apply(&self, input: &EvaluatedInput, truncation: Option<u64>) -> Vec<Tuple> {
if index_fn_supports_lookup(&input.index_fn) {
let mis = extrema_multi_indices(&input.index_fn, truncation);
mis.into_iter()
.filter_map(|mi| multi_index_to_flat(&input.index_fn, &mi))
.filter_map(|flat| input.tuples.get(flat).cloned())
.collect()
} else {
naive_extrema_prefix(&input.tuples, truncation)
}
}
}
fn naive_extrema_prefix(input: &[Tuple], truncation: Option<u64>) -> Vec<Tuple> {
if input.is_empty() {
return Vec::new();
}
let n = truncation
.unwrap_or(input.len() as u64)
.min(input.len() as u64);
if n == 0 {
return Vec::new();
}
if n == 1 {
return vec![input[0].clone()];
}
let mut out = Vec::with_capacity(n as usize);
out.push(input[0].clone());
out.push(input[input.len() - 1].clone());
for item in input.iter().take(input.len() - 1).skip(1) {
if (out.len() as u64) >= n {
break;
}
out.push(item.clone());
}
out
}
pub(crate) fn extrema_multi_indices(
idx: &IndexFn,
truncation: Option<u64>,
) -> Vec<MultiIndex> {
if index_fn_dim(idx) == 0 {
return Vec::new();
}
let axis_sizes: 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![axis_sizes.iter().copied().max().unwrap_or(0)]
}
IndexFn::Concatenation { segment_sizes } => {
vec![segment_sizes.iter().copied().sum()]
}
};
extrema_strata(&axis_sizes, truncation)
}
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> {
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)));
take_n_strata(scored, truncation)
}
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 extrema_3x3x3_strata_match_srd_18d_example() {
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]]);
}
}