Skip to main content

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}