polydat_core/iteration/comprehension/strategies/
shells.rs1use super::{
29 MultiIndex, Selection, Strategy, capped, index_fn_size, index_fn_supports_lookup,
30 lex::lex_multi_indices,
31};
32use crate::iteration::comprehension::metadata::IndexFn;
33use crate::iteration::comprehension::strategy::StrategyName;
34
35pub struct Shells;
37
38impl Strategy for Shells {
39 fn name(&self) -> StrategyName {
40 StrategyName::Shells
41 }
42
43 fn selects_from_shape(&self) -> bool {
45 true
46 }
47
48 fn accepts_input(&self, idx: Option<&IndexFn>) -> bool {
49 match idx {
50 None => false,
51 Some(i) => !i.has_continuous_axis(),
52 }
53 }
54
55 fn has_closed_form_for(&self, idx: &IndexFn) -> bool {
56 matches!(idx, IndexFn::Lattice { .. })
57 }
58
59 fn select(
60 &self,
61 index_fn: &IndexFn,
62 cardinality: u64,
63 truncation: Option<u64>,
64 _seed: Option<u64>,
65 ) -> Selection {
66 if index_fn_supports_lookup(index_fn) {
67 let mis = shells_multi_indices(index_fn, truncation);
68 Selection::from_multi_indices(index_fn, mis, cardinality)
69 } else {
70 Selection::Prefix(capped(truncation, cardinality))
71 }
72 }
73}
74
75pub(crate) fn shells_multi_indices(idx: &IndexFn, truncation: Option<u64>) -> Vec<MultiIndex> {
76 let total = index_fn_size(idx);
77 let n = match truncation {
78 Some(t) => t.min(total),
79 None => total,
80 };
81
82 let axis_sizes = match idx {
83 IndexFn::Lattice { axis_sizes } => axis_sizes.clone(),
84 _ => return lex_multi_indices(idx, truncation),
85 };
86
87 if axis_sizes.is_empty() {
88 return Vec::new();
89 }
90
91 let centers: Vec<f64> = axis_sizes.iter().map(|s| (*s as f64 - 1.0) / 2.0).collect();
92
93 let mut buckets: std::collections::BTreeMap<i64, Vec<MultiIndex>> =
96 std::collections::BTreeMap::new();
97
98 enumerate_all(
99 &axis_sizes,
100 &mut Vec::with_capacity(axis_sizes.len()),
101 &mut |mi| {
102 let r = chebyshev_distance(mi, ¢ers);
103 let bucket_key = (r * 2.0).round() as i64;
104 buckets.entry(bucket_key).or_default().push(mi.clone());
105 },
106 );
107
108 let mut out: Vec<MultiIndex> = Vec::with_capacity(n as usize);
109 for (_key, mut shell) in buckets.into_iter().rev() {
110 shell.sort();
111 for mi in shell {
112 if out.len() as u64 >= n {
113 return out;
114 }
115 out.push(mi);
116 }
117 }
118 out
119}
120
121fn chebyshev_distance(mi: &[u64], center: &[f64]) -> f64 {
122 mi.iter()
123 .zip(center.iter())
124 .map(|(c, ctr)| ((*c as f64) - ctr).abs())
125 .fold(0.0f64, f64::max)
126}
127
128fn enumerate_all(
129 axis_sizes: &[u64],
130 current: &mut Vec<u64>,
131 callback: &mut dyn FnMut(&MultiIndex),
132) {
133 if current.len() == axis_sizes.len() {
134 callback(current);
135 return;
136 }
137 let size = axis_sizes[current.len()];
138 for v in 0..size {
139 current.push(v);
140 enumerate_all(axis_sizes, current, callback);
141 current.pop();
142 }
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148
149 #[test]
150 fn shells_3x3() {
151 let idx = IndexFn::Lattice {
152 axis_sizes: vec![3, 3],
153 };
154 let out = shells_multi_indices(&idx, None);
155 assert_eq!(out.len(), 9);
156 assert_eq!(out[8], vec![1, 1]);
158 }
159
160 #[test]
161 fn shells_outermost_first_for_5x5() {
162 let idx = IndexFn::Lattice {
163 axis_sizes: vec![5, 5],
164 };
165 let out = shells_multi_indices(&idx, Some(4));
166 for mi in &out {
167 let chebyshev = chebyshev_distance(mi, &[2.0, 2.0]);
168 assert!(
169 (chebyshev - 2.0).abs() < 1e-9,
170 "expected distance 2.0, got {chebyshev:?}"
171 );
172 }
173 }
174
175 #[test]
176 fn shells_are_chebyshev_strata_outermost_first() {
177 let idx = IndexFn::Lattice {
182 axis_sizes: vec![5, 5],
183 };
184 let out = shells_multi_indices(&idx, None);
185 assert_eq!(out.len(), 25);
186 let dists: Vec<f64> = out
187 .iter()
188 .map(|mi| chebyshev_distance(mi, &[2.0, 2.0]))
189 .collect();
190 assert!(
191 dists.windows(2).all(|w| w[0] >= w[1] - 1e-9),
192 "shell distances not outermost-first: {dists:?}"
193 );
194 let r2 = dists.iter().filter(|d| (**d - 2.0).abs() < 1e-9).count();
197 let r1 = dists.iter().filter(|d| (**d - 1.0).abs() < 1e-9).count();
198 let r0 = dists.iter().filter(|d| **d < 1e-9).count();
199 assert_eq!((r2, r1, r0), (16, 8, 1));
200 }
201
202 #[test]
203 fn rejects_continuous() {
204 use crate::iteration::comprehension::cardinality::{Interval, ProductMeasure};
205 let cont = IndexFn::Continuous {
206 intervals: vec![Interval::closed(0.0, 1.0)],
207 measure: ProductMeasure::Uniform,
208 };
209 assert!(!Shells.accepts_input(Some(&cont)));
210 }
211}