polydat_core/iteration/comprehension/strategies/mod.rs
1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Strategy implementations — comprehension_forms.md §3.6, §10.2 R2,
5//! §10.7.8.
6//!
7//! ## Selection, then lookup
8//!
9//! A strategy's order is a function of its input's shape alone: the
10//! input's `IndexFn`, its tuple count, the truncation, and the seed.
11//! [`Strategy::select`] computes that order as a [`Selection`] of
12//! positions into the input without seeing a tuple, which is what
13//! lets an index-addressed evaluator choose `order halton/100`'s
14//! tuples from a large product and compute only those 100 (§10.2
15//! R2). [`Strategy::select_surviving`] is the selection over a
16//! filter's input of which only some positions pass (§5 V5).
17//! [`Strategy::apply`] is the selection looked up against an
18//! [`EvaluatedInput`]'s materialized tuples.
19//!
20//! Per §10.7.8 this is the **strategy invocation contract**: V4
21//! fires at invocation time against the input's evaluated
22//! `index_fn`, however the input source was authored (literal,
23//! range, context-free generator, or workload-param).
24//!
25//! Each strategy module holds a closed-form path over an `IndexFn`
26//! that supports lookup and a fallback over a one-axis position
27//! range; both produce positions.
28//!
29//! Strategies are selected by [`StrategyName`]; [`for_name`]
30//! dispatches a strategy name to its boxed [`Strategy`] impl.
31
32use super::ast::Comprehension;
33use super::metadata::{IndexFn, cycle_length};
34use super::strategy::StrategyName;
35
36pub mod antidiagonal;
37pub mod diagonal;
38pub mod extrema;
39pub mod halton;
40pub mod lex;
41pub mod lhs;
42pub mod prng;
43pub mod reverse_lex;
44pub mod shells;
45pub mod shuffle;
46pub mod sobol;
47
48/// A multi-coordinate index. Each component is the per-axis
49/// position in the input's index space. Length equals the
50/// input's dimensionality (1 for `Lockstep` / `Modular` /
51/// `Concatenation`; N for `Lattice` / `Continuous` /
52/// `Hybrid`).
53///
54/// `MultiIndex` is the indexed-form output type. The R2 IR
55/// opcode emitted by the IR compiler consumes these and resolves
56/// each through the input's `IndexFn` to dispense the actual
57/// tuple.
58pub type MultiIndex = Vec<u64>;
59
60/// A named-tuple value. Subset of the polydat `Value` set that
61/// is the strategy layer's currency; the runtime walker
62/// converts `Value`s to it before `apply` and maps results
63/// back. For the strategy module in isolation, this
64/// lightweight type lets tests run without pulling in the
65/// broader runtime.
66#[derive(Debug, Clone, PartialEq)]
67pub struct Tuple {
68 /// The tuple's `(name, value)` pairs, in shape order.
69 pub bindings: Vec<(String, TupleValue)>,
70}
71
72/// Subset of polydat's `Value` enum. `TupleValue` is the
73/// strategy layer's currency; the runtime walker converts
74/// `Value`s to it before `apply` and maps results back.
75#[derive(Debug, Clone, PartialEq)]
76pub enum TupleValue {
77 /// An unsigned integer.
78 U64(u64),
79 /// A signed integer.
80 I64(i64),
81 /// A float.
82 F64(f64),
83 /// A string.
84 Str(String),
85 /// A boolean.
86 Bool(bool),
87}
88
89impl Tuple {
90 /// An empty tuple.
91 pub fn new() -> Self {
92 Self {
93 bindings: Vec::new(),
94 }
95 }
96
97 /// The tuple with one more binding.
98 pub fn with<K: Into<String>>(mut self, key: K, value: TupleValue) -> Self {
99 self.bindings.push((key.into(), value));
100 self
101 }
102}
103
104impl Default for Tuple {
105 fn default() -> Self {
106 Self::new()
107 }
108}
109
110/// The materialized input to a strategy at invocation time
111/// (comprehension_forms.md §10.7.8).
112///
113/// `tuples` are the input stream's tuples in source order (the
114/// natural enumeration of the upstream comprehension subtree).
115/// `cardinality` matches `tuples.len() as u64`. `index_fn` is
116/// the addressing scheme the input actually satisfies —
117/// derived from observed shape for Generator /
118/// WorkloadParamList leaves via the [`crate::iteration::comprehension::eval_source`]
119/// layer, combined upward by the runtime walker per the
120/// propagation rules of comprehension_forms.md §10.7.2.
121pub struct EvaluatedInput {
122 /// The input's tuples, in source order.
123 pub tuples: Vec<Tuple>,
124 /// How many tuples: `tuples.len()`.
125 pub cardinality: u64,
126 /// The addressing scheme the input satisfies.
127 pub index_fn: IndexFn,
128}
129
130/// The positions a strategy emits, in emission order, as offsets
131/// into its input's natural enumeration.
132///
133/// A prefix and a reversal are held as their bounds; every other
134/// order is the list of positions it chose, one per emitted tuple.
135#[derive(Debug, Clone, PartialEq, Eq)]
136pub enum Selection {
137 /// Positions `0..n`.
138 Prefix(u64),
139 /// Positions `total - 1`, `total - 2`, …, `len` of them.
140 Reverse {
141 /// The input's tuple count.
142 total: u64,
143 /// How many positions are emitted.
144 len: u64,
145 },
146 /// The chosen positions, each below the input's tuple count.
147 Positions(Vec<u64>),
148}
149
150impl Selection {
151 /// How many positions the selection emits.
152 pub fn len(&self) -> u64 {
153 match self {
154 Selection::Prefix(n) => *n,
155 Selection::Reverse { len, .. } => *len,
156 Selection::Positions(p) => p.len() as u64,
157 }
158 }
159
160 /// Whether the selection emits nothing.
161 pub fn is_empty(&self) -> bool {
162 self.len() == 0
163 }
164
165 /// The input position emitted at `i`, or `None` past the end.
166 pub fn get(&self, i: u64) -> Option<u64> {
167 match self {
168 Selection::Prefix(n) => (i < *n).then_some(i),
169 Selection::Reverse { total, len } => (i < *len).then(|| total - 1 - i),
170 Selection::Positions(p) => usize::try_from(i).ok().and_then(|i| p.get(i).copied()),
171 }
172 }
173
174 /// The emitted positions, in order.
175 pub fn iter(&self) -> impl Iterator<Item = u64> + '_ {
176 (0..self.len()).filter_map(|i| self.get(i))
177 }
178
179 /// The positions of `multi_indices` over `idx`, keeping those
180 /// that land below `cardinality`.
181 pub(crate) fn from_multi_indices(
182 idx: &IndexFn,
183 multi_indices: Vec<MultiIndex>,
184 cardinality: u64,
185 ) -> Self {
186 Selection::Positions(
187 multi_indices
188 .into_iter()
189 .filter_map(|mi| multi_index_to_flat(idx, &mi))
190 .map(|flat| flat as u64)
191 .filter(|p| *p < cardinality)
192 .collect(),
193 )
194 }
195}
196
197/// The strategy invocation surface of comprehension_forms.md §10.7.8.
198///
199/// Implementations are stateless — every call to
200/// [`select`](Strategy::select) produces the same positions given the
201/// same inputs (deterministic). PRNG-based strategies (`Shuffle`,
202/// `Lhs`) derive their state from the authored seed, or a module
203/// constant when none is authored, plus the input length; no
204/// per-streamer seed is threaded.
205pub trait Strategy {
206 /// The strategy's name. Mirrors [`StrategyName`].
207 fn name(&self) -> StrategyName;
208
209 /// Whether the strategy selects from its input's shape rather than
210 /// from the sequence the input's tuples arrive in
211 /// (comprehension_forms.md §7.4 O1). A strategy that selects from
212 /// the shape places each tuple by its position in the input's index
213 /// space: it samples that space or walks its geometry. It chooses
214 /// the same tuples, in the same order, whatever permutation an
215 /// untruncated order applied to its input first, so that inner
216 /// order has no effect and is dropped (R7). A strategy that selects
217 /// from the sequence (a prefix, a reversal, a permutation of the
218 /// positions it is given) chooses differently after a permutation,
219 /// and both orders run.
220 fn selects_from_shape(&self) -> bool;
221
222 /// V4 input-shape check (comprehension_forms.md §3.6). `None` represents an
223 /// input with no closed-form index function; only `Lex`
224 /// accepts that. Concrete `IndexFn` variants are accepted
225 /// per the per-strategy rules in §3.6's table.
226 fn accepts_input(&self, idx: Option<&IndexFn>) -> bool;
227
228 /// R2 push-down eligibility (§10.2 R2). `true` if this
229 /// strategy has a closed-form multi-index rule over the given
230 /// input; otherwise [`select`](Strategy::select) orders the
231 /// input's positions as one axis.
232 fn has_closed_form_for(&self, idx: &IndexFn) -> bool;
233
234 /// The positions this strategy emits over an input of
235 /// `cardinality` tuples addressed by `index_fn`, cut to
236 /// `truncation`, under the authored `seed` (comprehension_forms.md
237 /// §3.6: a seeded strategy, `Shuffle` or `Lhs`, derives its state
238 /// from the seed and the input's structural identity, and from its
239 /// fixed default when `seed` is `None`; every other strategy
240 /// ignores it).
241 ///
242 /// The selection reads no tuple, so a caller that can compute the
243 /// tuple at a position computes only the selected ones. V4 is the
244 /// caller's responsibility: call `accepts_input` first.
245 fn select(
246 &self,
247 index_fn: &IndexFn,
248 cardinality: u64,
249 truncation: Option<u64>,
250 seed: Option<u64>,
251 ) -> Selection;
252
253 /// The positions this strategy emits over an input of which only
254 /// the positions in `survivors` (ascending) pass a filter
255 /// (comprehension_forms.md §5 V5): the strategy selects from the
256 /// input's whole index space, keeps the survivors in the order it
257 /// emits them, and applies its truncation to them. The positions
258 /// are the survivors' original positions in the input, so
259 /// `order(filter(c, p), halton, n)` yields `n` survivors whenever
260 /// at least `n` exist, and when every tuple survives the selection
261 /// is [`select`](Strategy::select)'s.
262 ///
263 /// Under a truncation `n` the strategy selects `n` positions, then
264 /// twice as many, and so on up to the whole input, until `n`
265 /// survivors are among them; at the whole input, survivors it does
266 /// not reach follow in ascending order. Without a truncation every
267 /// survivor is kept. A strategy whose truncation counts something
268 /// other than positions (`Extrema`'s strata) overrides this.
269 fn select_surviving(
270 &self,
271 index_fn: &IndexFn,
272 cardinality: u64,
273 truncation: Option<u64>,
274 seed: Option<u64>,
275 survivors: &[u64],
276 ) -> Selection {
277 surviving_in_rank(
278 &|count| self.select(index_fn, cardinality, count, seed),
279 cardinality,
280 truncation,
281 survivors,
282 )
283 }
284
285 /// Apply this strategy to the given input: its
286 /// [`select`](Strategy::select)ion looked up against
287 /// `input.tuples`.
288 ///
289 /// V4 is the caller's responsibility — call
290 /// `accepts_input(Some(&input.index_fn))` before `apply`
291 /// to fire V4 at strategy-invocation time per §10.7.8.
292 fn apply(&self, input: &EvaluatedInput, truncation: Option<u64>) -> Vec<Tuple> {
293 self.apply_seeded(input, truncation, None)
294 }
295
296 /// [`apply`](Strategy::apply) under an authored seed.
297 fn apply_seeded(
298 &self,
299 input: &EvaluatedInput,
300 truncation: Option<u64>,
301 seed: Option<u64>,
302 ) -> Vec<Tuple> {
303 self.select(&input.index_fn, input.tuples.len() as u64, truncation, seed)
304 .iter()
305 .filter_map(|p| input.tuples.get(p as usize).cloned())
306 .collect()
307 }
308}
309
310/// The first `truncation` of `survivors` (ascending positions) in the
311/// order `select` emits them, as [`Strategy::select_surviving`]
312/// describes: `select(Some(k))` for `k` from the truncation doubling up
313/// to `cardinality`, or `select(None)` without a truncation, followed at
314/// the whole input by the survivors it does not reach.
315pub(crate) fn surviving_in_rank(
316 select: &dyn Fn(Option<u64>) -> Selection,
317 cardinality: u64,
318 truncation: Option<u64>,
319 survivors: &[u64],
320) -> Selection {
321 let want = capped(truncation, survivors.len() as u64) as usize;
322 if want == 0 {
323 return Selection::Positions(Vec::new());
324 }
325 let mut count = truncation.map(|t| t.min(cardinality));
326 loop {
327 let whole = count.is_none_or(|k| k >= cardinality);
328 let selected = select(count);
329 let reached = selected
330 .iter()
331 .filter(|p| survivors.binary_search(p).is_ok());
332 let rest = survivors.iter().copied().filter(|_| whole);
333 let mut taken = std::collections::HashSet::with_capacity(want);
334 let mut out = Vec::with_capacity(want);
335 for p in reached.chain(rest) {
336 if out.len() == want {
337 break;
338 }
339 if taken.insert(p) {
340 out.push(p);
341 }
342 }
343 if out.len() == want || whole {
344 return Selection::Positions(out);
345 }
346 count = count.map(|k| k.saturating_mul(2).min(cardinality));
347 }
348}
349
350/// `n` capped at `total`, or `total` when there is no cap.
351pub(crate) fn capped(truncation: Option<u64>, total: u64) -> u64 {
352 truncation.map_or(total, |t| t.min(total))
353}
354
355/// Dispatch a [`StrategyName`] to its concrete [`Strategy`]
356/// implementation. The returned trait object is stateless;
357/// callers can hold a single instance per strategy name for
358/// the life of the process if desired.
359pub fn for_name(name: StrategyName) -> Box<dyn Strategy + Send + Sync> {
360 match name {
361 StrategyName::Lex => Box::new(lex::Lex),
362 StrategyName::ReverseLex => Box::new(reverse_lex::ReverseLex),
363 StrategyName::Shuffle => Box::new(shuffle::Shuffle),
364 StrategyName::Halton => Box::new(halton::Halton),
365 StrategyName::Sobol => Box::new(sobol::Sobol),
366 StrategyName::Lhs => Box::new(lhs::Lhs),
367 StrategyName::Extrema => Box::new(extrema::Extrema),
368 StrategyName::Shells => Box::new(shells::Shells),
369 StrategyName::Diagonal => Box::new(diagonal::Diagonal),
370 StrategyName::Antidiagonal => Box::new(antidiagonal::Antidiagonal),
371 }
372}
373
374/// The comprehension an order under `strategy` selects from, given its
375/// operand `child` (comprehension_forms.md §7.4 O1). A strategy that
376/// selects from its input's shape ([`Strategy::selects_from_shape`])
377/// reads through every untruncated order directly under it, since such
378/// an order only permutes the tuples of the shape beneath it; any other
379/// strategy selects from `child` itself.
380pub fn shape_input(child: &Comprehension, strategy: StrategyName) -> &Comprehension {
381 if !for_name(strategy).selects_from_shape() {
382 return child;
383 }
384 let mut input = child;
385 while let Comprehension::Order {
386 child,
387 truncation: None,
388 ..
389 } = input
390 {
391 input = child;
392 }
393 input
394}
395
396/// The filter a non-`Lex` order under `strategy` ranks the survivors of
397/// (comprehension_forms.md §5 V5), given its operand `child`: its
398/// predicate and the input the survivors' positions are taken in. The
399/// order selects from [`shape_input`]; when that is a filter, the
400/// survivors are ranked by their positions in the filter's input, which a
401/// strategy that selects from the shape reads through the untruncated
402/// orders of, as it does above the filter (§7.4 O1): such an order only
403/// permutes the tuples the predicate tests. `None` when the order ranks
404/// no filter.
405pub fn ranked_filter(
406 child: &Comprehension,
407 strategy: StrategyName,
408) -> Option<(&Comprehension, &str)> {
409 if strategy == StrategyName::Lex {
410 return None;
411 }
412 match shape_input(child, strategy) {
413 Comprehension::Filter { child, predicate } => {
414 Some((shape_input(child, strategy), predicate.as_str()))
415 }
416 _ => None,
417 }
418}
419
420/// Resolve a [`MultiIndex`] to a flat position in the
421/// input's tuple list, given the input's [`IndexFn`].
422///
423/// The flat position matches the natural enumeration order
424/// the runtime walker produces:
425///
426/// - `Lattice { axis_sizes: [s0, s1, …, sN-1] }` — row-major
427/// over the axes: `flat = i0 * s1 * s2 * … + i1 * s2 * … + … + iN-1`.
428/// This matches the runtime walker's cartesian enumeration
429/// (head axis varies slowest, tail nested).
430/// - `Lockstep { length }` — one-axis identity:
431/// `flat = mi[0]`.
432/// - `Modular { axis_sizes }` — one-axis identity over `max(axis_sizes)`:
433/// `flat = mi[0]`.
434/// - `Concatenation { segment_sizes }` — one-axis identity
435/// over `Σ segment_sizes`: `flat = mi[0]`.
436/// - `Continuous` / `Hybrid` — `None`; these inputs have no
437/// pre-materialized tuple list (the strategy's multi-indices
438/// are quantiles, not lookups).
439///
440/// Returns `None` for out-of-range positions or dimension
441/// mismatches.
442pub fn multi_index_to_flat(idx: &IndexFn, mi: &MultiIndex) -> Option<usize> {
443 match idx {
444 IndexFn::Lattice { axis_sizes } => {
445 if mi.len() != axis_sizes.len() {
446 return None;
447 }
448 let mut flat: u64 = 0;
449 let mut stride: u64 = 1;
450 for i in (0..axis_sizes.len()).rev() {
451 let pos = mi[i];
452 let size = axis_sizes[i];
453 if pos >= size {
454 return None;
455 }
456 flat = flat.checked_add(pos.checked_mul(stride)?)?;
457 stride = stride.checked_mul(size)?;
458 }
459 Some(flat as usize)
460 }
461 IndexFn::Lockstep { length } => {
462 if mi.len() != 1 || mi[0] >= *length {
463 return None;
464 }
465 Some(mi[0] as usize)
466 }
467 IndexFn::Modular { axis_sizes } => {
468 if mi.len() != 1 || mi[0] >= cycle_length(axis_sizes) {
469 return None;
470 }
471 Some(mi[0] as usize)
472 }
473 IndexFn::Concatenation { segment_sizes } => {
474 let total: u64 = segment_sizes.iter().copied().sum();
475 if mi.len() != 1 || mi[0] >= total {
476 return None;
477 }
478 Some(mi[0] as usize)
479 }
480 IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => None,
481 }
482}
483
484/// `true` when [`multi_index_to_flat`] returns a usable
485/// position for in-range multi-indices over this `IndexFn`.
486/// `false` for `Continuous` / `Hybrid` where the indexed
487/// strategy emits quantiles, not lookups.
488pub fn index_fn_supports_lookup(idx: &IndexFn) -> bool {
489 !matches!(idx, IndexFn::Continuous { .. } | IndexFn::Hybrid { .. })
490}
491
492/// Cardinality of an `IndexFn`. Used by strategies to size
493/// their output when no truncation is specified. Mirrors the
494/// helper in `metadata.rs` but lives here to avoid a circular
495/// dependency.
496pub(crate) fn index_fn_size(idx: &IndexFn) -> u64 {
497 match idx {
498 IndexFn::Lattice { axis_sizes } => axis_sizes
499 .iter()
500 .copied()
501 .fold(1u64, |a, b| a.saturating_mul(b)),
502 IndexFn::Lockstep { length } => *length,
503 IndexFn::Modular { axis_sizes } => cycle_length(axis_sizes),
504 IndexFn::Concatenation { segment_sizes } => segment_sizes
505 .iter()
506 .copied()
507 .fold(0u64, |a, b| a.saturating_add(b)),
508 IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => 0,
509 }
510}
511
512/// Lattice dimensionality of an `IndexFn`. Used by strategies
513/// that branch on dimensionality (Extrema's corner count,
514/// Lhs's per-axis stratification).
515pub(crate) fn index_fn_dim(idx: &IndexFn) -> usize {
516 match idx {
517 IndexFn::Lattice { axis_sizes } => axis_sizes.len(),
518 IndexFn::Continuous { intervals, .. } => intervals.len(),
519 IndexFn::Hybrid {
520 discrete_axes,
521 continuous_axes,
522 ..
523 } => discrete_axes.len() + continuous_axes.len(),
524 // A zip and a union are one axis of positions, which is what
525 // `multi_index_to_flat` reads from them.
526 IndexFn::Lockstep { .. } | IndexFn::Modular { .. } | IndexFn::Concatenation { .. } => 1,
527 }
528}
529
530#[cfg(test)]
531mod tests {
532 use super::*;
533
534 #[test]
535 fn for_name_dispatches_to_correct_strategy() {
536 assert_eq!(for_name(StrategyName::Lex).name(), StrategyName::Lex);
537 assert_eq!(for_name(StrategyName::Halton).name(), StrategyName::Halton);
538 assert_eq!(
539 for_name(StrategyName::Extrema).name(),
540 StrategyName::Extrema
541 );
542 }
543
544 #[test]
545 fn index_fn_size_lattice() {
546 let idx = IndexFn::Lattice {
547 axis_sizes: vec![3, 4, 5],
548 };
549 assert_eq!(index_fn_size(&idx), 60);
550 }
551
552 #[test]
553 fn index_fn_size_concatenation() {
554 let idx = IndexFn::Concatenation {
555 segment_sizes: vec![10, 20, 30],
556 };
557 assert_eq!(index_fn_size(&idx), 60);
558 }
559
560 #[test]
561 fn index_fn_dim_classifies_correctly() {
562 assert_eq!(
563 index_fn_dim(&IndexFn::Lattice {
564 axis_sizes: vec![3, 4]
565 }),
566 2
567 );
568 assert_eq!(index_fn_dim(&IndexFn::Lockstep { length: 10 }), 1);
569 assert_eq!(
570 index_fn_dim(&IndexFn::Concatenation {
571 segment_sizes: vec![1, 2, 3]
572 }),
573 1
574 );
575 }
576
577 /// Every strategy over a zip or a union emits positions within the
578 /// input, one axis as long as the input, and never fails.
579 #[test]
580 fn one_axis_inputs_select_within_their_length() {
581 let inputs = [
582 IndexFn::Modular {
583 axis_sizes: vec![2, 7, 3],
584 },
585 IndexFn::Concatenation {
586 segment_sizes: vec![2, 3, 4],
587 },
588 IndexFn::Lockstep { length: 9 },
589 ];
590 for idx in &inputs {
591 let total = index_fn_size(idx);
592 for name in [
593 StrategyName::Lex,
594 StrategyName::ReverseLex,
595 StrategyName::Diagonal,
596 StrategyName::Antidiagonal,
597 StrategyName::Extrema,
598 StrategyName::Shells,
599 StrategyName::Halton,
600 StrategyName::Sobol,
601 StrategyName::Lhs,
602 StrategyName::Shuffle,
603 ] {
604 let full: Vec<u64> = for_name(name)
605 .select(idx, total, None, None)
606 .iter()
607 .collect();
608 let mut sorted = full.clone();
609 sorted.sort_unstable();
610 sorted.dedup();
611 assert!(
612 full.iter().all(|p| *p < total),
613 "{name:?} over {idx:?}: {full:?}"
614 );
615 if !matches!(name, StrategyName::Halton | StrategyName::Sobol) {
616 assert_eq!(
617 sorted.len() as u64,
618 total,
619 "{name:?} over {idx:?} reaches every position: {full:?}"
620 );
621 }
622 }
623 }
624 }
625}