#[cfg(test)]
use std::sync::Arc;
use crate::ast::Value;
use crate::iteration::comprehension::ast::Comprehension;
use crate::iteration::comprehension::cardinality::{Interval, ProductMeasure};
use crate::iteration::comprehension::eval_source::{EvalContext, SourceEval};
use crate::iteration::comprehension::measure::AxisMeasure;
use crate::iteration::comprehension::metadata::{IndexFn, cycle_length};
use crate::iteration::comprehension::predicate::CompiledPredicate;
use crate::iteration::comprehension::source::Source;
use crate::iteration::comprehension::strategies::{Selection, ranked_filter, shape_input};
use crate::iteration::comprehension::strategy::StrategyName;
#[cfg(test)]
use crate::kernel::PolydatKernel;
use crate::kernel::interp::Lookup;
pub type RuntimeTuple = Vec<(String, Value)>;
struct EvaluatedNode {
tuples: Vec<RuntimeTuple>,
index_fn: Option<IndexFn>,
}
#[derive(Debug, Clone)]
pub enum RuntimeError {
SourceEval {
var: String,
source: String,
message: String,
},
FilterEval {
predicate: String,
message: String,
},
OrderEval {
strategy: StrategyName,
message: String,
},
StrategyRejectsInput {
strategy: StrategyName,
index_fn: Option<IndexFn>,
},
UnsupportedShape(String),
ZipLengthMismatch {
lengths: Vec<u64>,
},
}
impl std::fmt::Display for RuntimeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RuntimeError::SourceEval {
var,
source,
message,
} => {
write!(f, "for_each clause '{var} in {source}': {message}")
}
RuntimeError::FilterEval { predicate, message } => {
write!(f, "comprehension filter '{predicate}': {message}")
}
RuntimeError::OrderEval { strategy, message } => {
write!(f, "order strategy {strategy:?}: {message}")
}
RuntimeError::StrategyRejectsInput { strategy, index_fn } => write!(
f,
"order strategy {strategy:?} rejects input shape {index_fn:?} \
(V4: per-strategy IndexFn contract; see comprehension_forms.md §3.6's \
strategy table)"
),
RuntimeError::UnsupportedShape(msg) => write!(f, "{msg}"),
RuntimeError::ZipLengthMismatch { lengths } => {
write!(f, "zip strict: child lengths differ ({lengths:?})")
}
}
}
}
impl std::error::Error for RuntimeError {}
pub fn evaluate_for_iteration(
comp: &Comprehension,
scope: &dyn Lookup,
) -> Result<Vec<RuntimeTuple>, RuntimeError> {
evaluate_indexed(comp, scope).map(|t| t.to_vec())
}
pub fn evaluate_for_iteration_reported(
comp: &Comprehension,
scope: &dyn Lookup,
) -> Result<EvaluatedIteration, RuntimeError> {
let mut state = EvalState::new(comp, scope);
let (node, _) = state.index_node(comp, &[])?;
Ok(EvaluatedIteration {
tuples: IndexedTuples { node }.to_vec(),
clauses: state.yields,
})
}
pub fn evaluate_indexed(
comp: &Comprehension,
scope: &dyn Lookup,
) -> Result<IndexedTuples, RuntimeError> {
let mut state = EvalState::new(comp, scope);
let (node, _) = state.index_node(comp, &[])?;
Ok(IndexedTuples { node })
}
pub fn evaluate_for_iteration_materialized(
comp: &Comprehension,
scope: &dyn Lookup,
) -> Result<EvaluatedIteration, RuntimeError> {
let mut state = EvalState::new(comp, scope);
let tuples = state.evaluate_node(comp, &[])?.tuples;
Ok(EvaluatedIteration {
tuples,
clauses: state.yields,
})
}
#[derive(Debug, Clone)]
pub struct IndexedTuples {
node: Indexed,
}
impl IndexedTuples {
pub fn len(&self) -> u64 {
self.node.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn get(&self, i: u64) -> Option<RuntimeTuple> {
if i >= self.len() {
return None;
}
let mut out = RuntimeTuple::new();
self.node.append_at(i, &mut out);
Some(out)
}
pub fn iter(&self) -> impl Iterator<Item = RuntimeTuple> + '_ {
(0..self.len()).filter_map(|i| self.get(i))
}
pub fn to_vec(&self) -> Vec<RuntimeTuple> {
self.iter().collect()
}
}
#[derive(Debug, Clone)]
enum Indexed {
Tuples(Vec<RuntimeTuple>),
Clause { name: String, values: ClauseValues },
Product {
children: Vec<Indexed>,
lens: Vec<u64>,
len: u64,
},
Lockstep { children: Vec<Indexed>, len: u64 },
Cycle { children: Vec<Indexed>, len: u64 },
Concat { children: Vec<Indexed>, len: u64 },
Select {
child: Box<Indexed>,
selection: Selection,
},
}
#[derive(Debug, Clone)]
enum ClauseValues {
Range { lo: i64, step: i64, len: u64 },
List(Vec<Value>),
}
impl ClauseValues {
fn len(&self) -> u64 {
match self {
ClauseValues::Range { len, .. } => *len,
ClauseValues::List(values) => values.len() as u64,
}
}
fn at(&self, i: u64) -> Value {
match self {
ClauseValues::Range { lo, step, .. } => {
Value::U64((i128::from(*lo) + i128::from(i) * i128::from(*step)) as i64 as u64)
}
ClauseValues::List(values) => values[i as usize].clone(),
}
}
}
impl Indexed {
fn len(&self) -> u64 {
match self {
Indexed::Tuples(tuples) => tuples.len() as u64,
Indexed::Clause { values, .. } => values.len(),
Indexed::Product { len, .. }
| Indexed::Lockstep { len, .. }
| Indexed::Cycle { len, .. }
| Indexed::Concat { len, .. } => *len,
Indexed::Select { selection, .. } => selection.len(),
}
}
fn append_at(&self, i: u64, out: &mut RuntimeTuple) {
match self {
Indexed::Tuples(tuples) => out.extend(tuples[i as usize].iter().cloned()),
Indexed::Clause { name, values } => out.push((name.clone(), values.at(i))),
Indexed::Product { children, lens, .. } => {
let mut digits = vec![0u64; lens.len()];
let mut rest = i;
for (d, len) in digits.iter_mut().zip(lens).rev() {
*d = rest % len;
rest /= len;
}
for (child, d) in children.iter().zip(digits) {
child.append_at(d, out);
}
}
Indexed::Lockstep { children, .. } => {
for child in children {
child.append_at(i, out);
}
}
Indexed::Cycle { children, .. } => {
for child in children {
child.append_at(i % child.len(), out);
}
}
Indexed::Concat { children, .. } => {
let mut offset = i;
for child in children {
let len = child.len();
if offset < len {
child.append_at(offset, out);
return;
}
offset -= len;
}
}
Indexed::Select { child, selection } => {
if let Some(p) = selection.get(i) {
child.append_at(p, out);
}
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ClauseYield {
pub var: String,
pub source: Option<String>,
pub evaluations: usize,
pub values: usize,
}
#[derive(Debug, Clone)]
pub struct EvaluatedIteration {
pub tuples: Vec<RuntimeTuple>,
pub clauses: Vec<ClauseYield>,
}
struct EvalState<'a> {
scope: &'a dyn Lookup,
yields: Vec<ClauseYield>,
by_leaf: std::collections::HashMap<usize, usize>,
mult: usize,
}
impl<'a> EvalState<'a> {
fn new(comp: &Comprehension, scope: &'a dyn Lookup) -> Self {
let mut state = EvalState {
scope,
yields: Vec::new(),
by_leaf: std::collections::HashMap::new(),
mult: 1,
};
state.enumerate_leaves(comp);
state
}
fn enumerate_leaves(&mut self, node: &Comprehension) {
match node {
Comprehension::Clause { name, source } => {
self.by_leaf
.insert(std::ptr::from_ref(source) as usize, self.yields.len());
self.yields.push(ClauseYield {
var: name.clone(),
source: source.to_text(),
evaluations: 0,
values: 0,
});
}
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => {
for child in children {
self.enumerate_leaves(child);
}
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
self.enumerate_leaves(child);
}
}
}
fn record_yield(&mut self, source: &Source, values: usize) {
if let Some(&i) = self.by_leaf.get(&(std::ptr::from_ref(source) as usize)) {
self.yields[i].evaluations = self.yields[i].evaluations.saturating_add(self.mult);
self.yields[i].values = self.yields[i]
.values
.saturating_add(values.saturating_mul(self.mult));
}
}
}
impl EvalState<'_> {
fn evaluate_node(
&mut self,
node: &Comprehension,
prefix: &[(String, Value)],
) -> Result<EvaluatedNode, RuntimeError> {
match node {
Comprehension::Clause { name, source } => self.evaluate_clause(name, source, prefix),
Comprehension::Cartesian { children } => self.evaluate_cartesian(children, prefix),
Comprehension::Zip { children, mode } => self.evaluate_zip(children, *mode, prefix),
Comprehension::Union { children } => self.evaluate_union(children, prefix),
Comprehension::Filter { child, predicate } => {
let inner = self.evaluate_node(child, prefix)?;
self.apply_filter(inner, predicate)
}
Comprehension::Order {
child,
strategy,
truncation,
seed,
} => {
let child = shape_input(child, *strategy);
if has_continuous_axis(child) {
let sampled =
self.sample_space(child, prefix, *strategy, *truncation, *seed)?;
return Ok(EvaluatedNode {
index_fn: Some(selected(sampled.tuples.len() as u64)),
tuples: sampled.tuples,
});
}
if let Some((child, predicate)) = ranked_filter(child, *strategy) {
let inner = self.evaluate_node(child, prefix)?;
let predicate = CompiledPredicate::new(predicate);
let mut survivors = Vec::new();
for (p, tuple) in inner.tuples.iter().enumerate() {
if predicate.keeps(tuple, self.scope)? {
survivors.push(p as u64);
}
}
let selection = surviving_selection(
*strategy,
inner.index_fn.as_ref(),
inner.tuples.len() as u64,
*truncation,
*seed,
&survivors,
)?;
return Ok(EvaluatedNode {
tuples: selection
.iter()
.map(|p| inner.tuples[p as usize].clone())
.collect(),
index_fn: Some(selected(selection.len())),
});
}
let inner = self.evaluate_node(child, prefix)?;
self.apply_order(inner, *strategy, *truncation, *seed)
}
}
}
fn evaluate_source(
&mut self,
name: &str,
source: &Source,
prefix: &[(String, Value)],
) -> Result<crate::iteration::comprehension::eval_source::EvaluatedSource, RuntimeError> {
let ctx = EvalContext {
var_name: name,
scope: self.scope,
prefix,
};
let evaluated = source.evaluate(Some(&ctx)).map_err(|e| match e {
crate::iteration::comprehension::eval_source::EvalError::EvalFailed {
var,
source,
message,
} => RuntimeError::SourceEval {
var,
source,
message,
},
crate::iteration::comprehension::eval_source::EvalError::NeedsContext => {
RuntimeError::UnsupportedShape(format!(
"clause '{name}': source requires kernel context but evaluator \
provided none — internal bug in runtime walker"
))
}
})?;
self.record_yield(source, evaluated.values.len());
Ok(evaluated)
}
fn evaluate_clause(
&mut self,
name: &str,
source: &Source,
prefix: &[(String, Value)],
) -> Result<EvaluatedNode, RuntimeError> {
let evaluated = self.evaluate_source(name, source, prefix)?;
if evaluated.values.is_empty() {
return Ok(EvaluatedNode {
tuples: Vec::new(),
index_fn: Some(evaluated.index_fn),
});
}
let tuples: Vec<RuntimeTuple> = evaluated
.values
.into_iter()
.map(|v| vec![(name.to_string(), v)])
.collect();
Ok(EvaluatedNode {
tuples,
index_fn: Some(evaluated.index_fn),
})
}
fn evaluate_cartesian(
&mut self,
children: &[Comprehension],
prefix: &[(String, Value)],
) -> Result<EvaluatedNode, RuntimeError> {
if children.is_empty() {
return Ok(EvaluatedNode {
tuples: vec![Vec::new()],
index_fn: Some(IndexFn::Lattice {
axis_sizes: vec![1],
}),
});
}
let mut child_index_fns: Vec<(Option<IndexFn>, u64)> = Vec::with_capacity(children.len());
let mut dependent_observed = false;
let result_tuples = self.evaluate_cartesian_rec(
children.len(),
children,
prefix,
&mut child_index_fns,
&mut dependent_observed,
)?;
let combined = if dependent_observed {
None
} else {
let index_fns: Vec<Option<IndexFn>> =
child_index_fns.into_iter().map(|(idx, _)| idx).collect();
combine_cartesian_index_fn(&index_fns)
};
Ok(EvaluatedNode {
tuples: result_tuples,
index_fn: combined,
})
}
fn evaluate_cartesian_rec(
&mut self,
child_count: usize,
children: &[Comprehension],
prefix: &[(String, Value)],
child_index_fns: &mut Vec<(Option<IndexFn>, u64)>,
dependent_observed: &mut bool,
) -> Result<Vec<RuntimeTuple>, RuntimeError> {
if children.is_empty() {
return Ok(vec![Vec::new()]);
}
let (head, tail) = children.split_first().unwrap();
let head_eval = self.evaluate_node(head, prefix)?;
let head_axis_len = head_eval.tuples.len() as u64;
let depth = child_count - children.len();
match child_index_fns.get(depth) {
None => child_index_fns.push((head_eval.index_fn.clone(), head_axis_len)),
Some((_, first)) if *first != head_axis_len => *dependent_observed = true,
Some(_) => {}
}
if tail.is_empty() {
return Ok(head_eval.tuples);
}
let mut out = Vec::new();
for head_tuple in head_eval.tuples {
let mut extended_prefix: Vec<(String, Value)> = prefix.to_vec();
extended_prefix.extend(head_tuple.iter().cloned());
let tail_tuples = self.evaluate_cartesian_rec(
child_count,
tail,
&extended_prefix,
child_index_fns,
dependent_observed,
)?;
for tail_tuple in tail_tuples {
let mut merged = head_tuple.clone();
merged.extend(tail_tuple);
out.push(merged);
}
}
Ok(out)
}
fn evaluate_zip(
&mut self,
children: &[Comprehension],
mode: crate::iteration::comprehension::strategy::ZipMode,
prefix: &[(String, Value)],
) -> Result<EvaluatedNode, RuntimeError> {
use crate::iteration::comprehension::strategy::ZipMode;
if children.is_empty() {
return Ok(EvaluatedNode {
tuples: vec![Vec::new()],
index_fn: Some(IndexFn::Lockstep { length: 1 }),
});
}
let per_child: Vec<EvaluatedNode> = children
.iter()
.map(|c| self.evaluate_node(c, prefix))
.collect::<Result<_, _>>()?;
let lengths: Vec<usize> = per_child.iter().map(|n| n.tuples.len()).collect();
let iter_count = match mode {
ZipMode::Strict => {
let first = lengths.first().copied().unwrap_or(0);
if lengths.iter().any(|&n| n != first) {
return Err(RuntimeError::ZipLengthMismatch {
lengths: lengths.iter().map(|&n| n as u64).collect(),
});
}
first
}
ZipMode::Truncate => lengths.iter().copied().min().unwrap_or(0),
ZipMode::Cycle => {
let counts: Vec<u64> = lengths.iter().map(|&n| n as u64).collect();
cycle_length(&counts) as usize
}
};
let mut tuples = Vec::with_capacity(iter_count);
for i in 0..iter_count {
let mut bindings: RuntimeTuple = Vec::new();
for (child, &len) in per_child.iter().zip(lengths.iter()) {
let idx = match mode {
ZipMode::Cycle => i % len,
_ => i,
};
bindings.extend(child.tuples[idx].iter().cloned());
}
tuples.push(bindings);
}
let index_fn = match mode {
ZipMode::Strict | ZipMode::Truncate => Some(IndexFn::Lockstep {
length: iter_count as u64,
}),
ZipMode::Cycle => Some(IndexFn::Modular {
axis_sizes: lengths.iter().map(|n| *n as u64).collect(),
}),
};
Ok(EvaluatedNode { tuples, index_fn })
}
fn evaluate_union(
&mut self,
children: &[Comprehension],
prefix: &[(String, Value)],
) -> Result<EvaluatedNode, RuntimeError> {
let mut tuples = Vec::new();
let mut segment_sizes = Vec::with_capacity(children.len());
let mut all_segments_addressable = true;
for child in children {
let sub = self.evaluate_node(child, prefix)?;
segment_sizes.push(sub.tuples.len() as u64);
if sub.index_fn.is_none() {
all_segments_addressable = false;
}
tuples.extend(sub.tuples);
}
let index_fn = if all_segments_addressable {
Some(IndexFn::Concatenation { segment_sizes })
} else {
None
};
Ok(EvaluatedNode { tuples, index_fn })
}
fn apply_filter(
&mut self,
input: EvaluatedNode,
predicate: &str,
) -> Result<EvaluatedNode, RuntimeError> {
let predicate = CompiledPredicate::new(predicate);
let mut out = Vec::with_capacity(input.tuples.len());
for tuple in input.tuples {
if predicate.keeps(&tuple, self.scope)? {
out.push(tuple);
}
}
Ok(EvaluatedNode {
tuples: out,
index_fn: None,
})
}
fn sample_space(
&mut self,
child: &Comprehension,
prefix: &[(String, Value)],
strategy: StrategyName,
truncation: Option<u64>,
seed: Option<u64>,
) -> Result<EvaluatedNode, RuntimeError> {
let mut space = SampleSpace::default();
self.collect_sample_space(child, prefix, &mut space, &mut Vec::new())?;
let discrete_axes: Vec<u64> = space
.axes
.iter()
.filter_map(|a| match a {
SampleAxis::Discrete(tuples) => Some(tuples.len() as u64),
SampleAxis::Continuous { .. } => None,
})
.collect();
if discrete_axes.contains(&0) {
return Ok(EvaluatedNode {
tuples: Vec::new(),
index_fn: None,
});
}
let (intervals, measures): (Vec<Interval>, Vec<ProductMeasure>) = space
.axes
.iter()
.filter_map(|a| match a {
SampleAxis::Continuous {
interval, measure, ..
} => Some((
interval.clone(),
match measure {
AxisMeasure::Uniform => ProductMeasure::Uniform,
AxisMeasure::Named { name, .. } => ProductMeasure::Named(*name),
},
)),
SampleAxis::Discrete(_) => None,
})
.unzip();
let sequence = !matches!(strategy, StrategyName::Extrema);
let index_fn = if !sequence {
IndexFn::Lattice {
axis_sizes: space
.axes
.iter()
.map(|a| match a {
SampleAxis::Discrete(tuples) => tuples.len() as u64,
SampleAxis::Continuous { .. } => 2,
})
.collect(),
}
} else if discrete_axes.is_empty() {
IndexFn::Continuous {
intervals,
measure: ProductMeasure::Product(measures),
}
} else {
IndexFn::Hybrid {
discrete_axes,
continuous_axes: intervals,
measure: ProductMeasure::Product(measures),
}
};
let mut want = truncation;
let mut rounds = 0;
loop {
let multi_indices = draw_sample(&index_fn, strategy, want, seed)?;
let drawn = multi_indices.len() as u64;
let tuples = multi_indices
.iter()
.map(|mi| space.realize(mi, strategy))
.collect();
let mut node = EvaluatedNode {
tuples,
index_fn: None,
};
for predicate in &space.predicates {
node = self.apply_filter(node, predicate)?;
}
let (Some(n), Some(asked)) = (truncation, want) else {
return Ok(node);
};
let enough = node.tuples.len() as u64 >= n;
let exhausted = drawn < asked;
if !sequence {
return Ok(node);
}
if enough || exhausted || rounds >= SAMPLE_ROUNDS {
node.tuples.truncate(n as usize);
return Ok(node);
}
want = Some(asked.saturating_mul(2));
rounds += 1;
}
}
fn collect_sample_space(
&mut self,
c: &Comprehension,
prefix: &[(String, Value)],
space: &mut SampleSpace,
bound: &mut Vec<String>,
) -> Result<(), RuntimeError> {
let measure_error = |name: &str, message: String| RuntimeError::SourceEval {
var: name.to_string(),
source: "<continuous>".to_string(),
message,
};
match c {
Comprehension::Clause {
name,
source: Source::ContinuousInterval { interval, measure },
} => {
let measure =
AxisMeasure::from_product(measure, 0).map_err(|m| measure_error(name, m))?;
space.axes.push(SampleAxis::Continuous {
name: name.clone(),
interval: interval.clone(),
measure,
});
bound.push(name.clone());
}
Comprehension::Clause {
name,
source:
Source::Distribution {
distribution,
support,
params,
},
} => {
let measure = AxisMeasure::named(*distribution, params)
.map_err(|m| measure_error(name, m))?;
space.axes.push(SampleAxis::Continuous {
name: name.clone(),
interval: support.clone(),
measure,
});
bound.push(name.clone());
}
Comprehension::Clause { name, source } => {
let references = c.referenced_source_names();
if let Some(dep) = bound.iter().find(|b| references.contains(*b)) {
return Err(RuntimeError::UnsupportedShape(format!(
"clause '{name}' references '{dep}' beside a continuous axis; \
a sampled cartesian is independent (comprehension_forms.md §5, V4)"
)));
}
let node = self.evaluate_clause(name, source, prefix)?;
space.axes.push(SampleAxis::Discrete(node.tuples));
bound.push(name.clone());
}
Comprehension::Cartesian { children } => {
for child in children {
self.collect_sample_space(child, prefix, space, bound)?;
}
}
Comprehension::Filter { child, predicate } => {
self.collect_sample_space(child, prefix, space, bound)?;
space.predicates.push(predicate.clone());
}
Comprehension::Zip { .. }
| Comprehension::Union { .. }
| Comprehension::Order { .. } => {
let node = self.evaluate_node(c, prefix)?;
bound.extend(c.coordinate_names());
space.axes.push(SampleAxis::Discrete(node.tuples));
}
}
Ok(())
}
fn apply_order(
&mut self,
input: EvaluatedNode,
strategy: StrategyName,
truncation: Option<u64>,
seed: Option<u64>,
) -> Result<EvaluatedNode, RuntimeError> {
let selection = order_selection(
strategy,
input.index_fn.as_ref(),
input.tuples.len() as u64,
truncation,
seed,
)?;
let out = selection
.iter()
.map(|p| input.tuples[p as usize].clone())
.collect();
Ok(EvaluatedNode {
tuples: out,
index_fn: order_output(strategy, truncation, input.index_fn, selection.len()),
})
}
}
fn selected(len: u64) -> IndexFn {
IndexFn::Lattice {
axis_sizes: vec![len],
}
}
fn order_output(
strategy: StrategyName,
truncation: Option<u64>,
input: Option<IndexFn>,
len: u64,
) -> Option<IndexFn> {
match (strategy, truncation) {
(StrategyName::Lex, None) => input,
(StrategyName::Lex, Some(_)) => input.map(|_| selected(len)),
_ => Some(selected(len)),
}
}
impl EvalState<'_> {
fn index_node(
&mut self,
node: &Comprehension,
prefix: &[(String, Value)],
) -> Result<(Indexed, Option<IndexFn>), RuntimeError> {
match node {
Comprehension::Clause { name, source } => self.index_clause(name, source, prefix),
Comprehension::Cartesian { children } => self.index_cartesian(children, prefix),
Comprehension::Zip { children, mode } => self.index_zip(children, *mode, prefix),
Comprehension::Union { children } => self.index_union(children, prefix),
Comprehension::Filter { child, predicate } => {
let (inner, _) = self.index_node(child, prefix)?;
let predicate = CompiledPredicate::new(predicate);
let mut kept = Vec::new();
let mut tuple = RuntimeTuple::new();
for i in 0..inner.len() {
tuple.clear();
inner.append_at(i, &mut tuple);
if predicate.keeps(&tuple, self.scope)? {
kept.push(tuple.clone());
}
}
Ok((Indexed::Tuples(kept), None))
}
Comprehension::Order {
child,
strategy,
truncation,
seed,
} => {
let child = shape_input(child, *strategy);
if has_continuous_axis(child) {
let sampled =
self.sample_space(child, prefix, *strategy, *truncation, *seed)?;
let len = sampled.tuples.len() as u64;
return Ok((Indexed::Tuples(sampled.tuples), Some(selected(len))));
}
if let Some((child, predicate)) = ranked_filter(child, *strategy) {
let (inner, index_fn) = self.index_node(child, prefix)?;
let predicate = CompiledPredicate::new(predicate);
let mut survivors = Vec::new();
let mut tuple = RuntimeTuple::new();
for p in 0..inner.len() {
tuple.clear();
inner.append_at(p, &mut tuple);
if predicate.keeps(&tuple, self.scope)? {
survivors.push(p);
}
}
let selection = surviving_selection(
*strategy,
index_fn.as_ref(),
inner.len(),
*truncation,
*seed,
&survivors,
)?;
let len = selection.len();
return Ok((
Indexed::Select {
child: Box::new(inner),
selection,
},
Some(selected(len)),
));
}
let (inner, index_fn) = self.index_node(child, prefix)?;
let selection = order_selection(
*strategy,
index_fn.as_ref(),
inner.len(),
*truncation,
*seed,
)?;
let len = selection.len();
Ok((
Indexed::Select {
child: Box::new(inner),
selection,
},
order_output(*strategy, *truncation, index_fn, len),
))
}
}
}
fn index_clause(
&mut self,
name: &str,
source: &Source,
prefix: &[(String, Value)],
) -> Result<(Indexed, Option<IndexFn>), RuntimeError> {
let values = match source {
Source::IntRange { lo, hi, step } => {
let step = (*step).max(1);
let len = if hi <= lo {
0
} else {
((i128::from(*hi) - i128::from(*lo)) as u128).div_ceil(step as u128) as u64
};
self.record_yield(source, len as usize);
ClauseValues::Range { lo: *lo, step, len }
}
_ => {
let evaluated = self.evaluate_source(name, source, prefix)?;
let index_fn = evaluated.index_fn;
return Ok((
Indexed::Clause {
name: name.to_string(),
values: ClauseValues::List(evaluated.values),
},
Some(index_fn),
));
}
};
let len = values.len();
Ok((
Indexed::Clause {
name: name.to_string(),
values,
},
Some(IndexFn::Lattice {
axis_sizes: vec![len],
}),
))
}
fn index_cartesian(
&mut self,
children: &[Comprehension],
prefix: &[(String, Value)],
) -> Result<(Indexed, Option<IndexFn>), RuntimeError> {
if children.is_empty() || references_an_earlier_axis(children) {
let node = self.evaluate_cartesian(children, prefix)?;
return Ok((Indexed::Tuples(node.tuples), node.index_fn));
}
let base = self.mult;
let mut parts = Vec::with_capacity(children.len());
let mut index_fns = Vec::with_capacity(children.len());
let mut lens = Vec::with_capacity(children.len());
let mut len: u64 = 1;
for child in children {
let evaluated = self.index_node(child, prefix);
let (part, index_fn) = match evaluated {
Ok(done) => done,
Err(e) => {
self.mult = base;
return Err(e);
}
};
let part_len = part.len();
parts.push(part);
index_fns.push(index_fn);
lens.push(part_len);
len = match len.checked_mul(part_len) {
Some(n) => n,
None => {
self.mult = base;
return Err(RuntimeError::UnsupportedShape(format!(
"cartesian of {lens:?} tuples exceeds 2^64"
)));
}
};
if part_len == 0 {
break;
}
self.mult = self
.mult
.saturating_mul(usize::try_from(part_len).unwrap_or(usize::MAX));
}
self.mult = base;
let index_fn = combine_cartesian_index_fn(&index_fns);
if len == 0 {
return Ok((Indexed::Tuples(Vec::new()), index_fn));
}
Ok((
Indexed::Product {
children: parts,
lens,
len,
},
index_fn,
))
}
fn index_zip(
&mut self,
children: &[Comprehension],
mode: crate::iteration::comprehension::strategy::ZipMode,
prefix: &[(String, Value)],
) -> Result<(Indexed, Option<IndexFn>), RuntimeError> {
use crate::iteration::comprehension::strategy::ZipMode;
if children.is_empty() {
let node = self.evaluate_zip(children, mode, prefix)?;
return Ok((Indexed::Tuples(node.tuples), node.index_fn));
}
let mut parts = Vec::with_capacity(children.len());
for child in children {
parts.push(self.index_node(child, prefix)?.0);
}
let lengths: Vec<u64> = parts.iter().map(Indexed::len).collect();
let len = match mode {
ZipMode::Strict => {
let first = lengths[0];
if lengths.iter().any(|&n| n != first) {
return Err(RuntimeError::ZipLengthMismatch { lengths });
}
first
}
ZipMode::Truncate => lengths.iter().copied().min().unwrap_or(0),
ZipMode::Cycle => cycle_length(&lengths),
};
Ok(match mode {
ZipMode::Strict | ZipMode::Truncate => (
Indexed::Lockstep {
children: parts,
len,
},
Some(IndexFn::Lockstep { length: len }),
),
ZipMode::Cycle => (
Indexed::Cycle {
children: parts,
len,
},
Some(IndexFn::Modular {
axis_sizes: lengths,
}),
),
})
}
fn index_union(
&mut self,
children: &[Comprehension],
prefix: &[(String, Value)],
) -> Result<(Indexed, Option<IndexFn>), RuntimeError> {
let mut parts = Vec::with_capacity(children.len());
let mut segment_sizes = Vec::with_capacity(children.len());
let mut all_segments_addressable = true;
for child in children {
let (part, index_fn) = self.index_node(child, prefix)?;
all_segments_addressable &= index_fn.is_some();
segment_sizes.push(part.len());
parts.push(part);
}
let len = segment_sizes
.iter()
.try_fold(0u64, |acc, n| acc.checked_add(*n))
.ok_or_else(|| {
RuntimeError::UnsupportedShape(format!(
"union of {segment_sizes:?} tuples exceeds 2^64"
))
})?;
let index_fn = all_segments_addressable.then_some(IndexFn::Concatenation { segment_sizes });
Ok((
Indexed::Concat {
children: parts,
len,
},
index_fn,
))
}
}
fn references_an_earlier_axis(children: &[Comprehension]) -> bool {
let mut bound: std::collections::BTreeSet<String> = std::collections::BTreeSet::new();
for child in children {
if child
.referenced_source_names()
.iter()
.any(|n| bound.contains(n))
{
return true;
}
collect_clause_names(child, &mut bound);
}
false
}
fn collect_clause_names(c: &Comprehension, out: &mut std::collections::BTreeSet<String>) {
match c {
Comprehension::Clause { name, .. } => {
out.insert(name.clone());
}
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => {
for child in children {
collect_clause_names(child, out);
}
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
collect_clause_names(child, out);
}
}
}
fn order_selection(
strategy: StrategyName,
index_fn: Option<&IndexFn>,
cardinality: u64,
truncation: Option<u64>,
seed: Option<u64>,
) -> Result<Selection, RuntimeError> {
let dispatch = crate::iteration::comprehension::strategies::for_name(strategy);
if !dispatch.accepts_input(index_fn) {
return Err(RuntimeError::StrategyRejectsInput {
strategy,
index_fn: index_fn.cloned(),
});
}
let fallback;
let index_fn = match index_fn {
Some(idx) => idx,
None => {
fallback = IndexFn::Lattice {
axis_sizes: vec![cardinality],
};
&fallback
}
};
Ok(dispatch.select(index_fn, cardinality, truncation, seed))
}
fn surviving_selection(
strategy: StrategyName,
index_fn: Option<&IndexFn>,
cardinality: u64,
truncation: Option<u64>,
seed: Option<u64>,
survivors: &[u64],
) -> Result<Selection, RuntimeError> {
let dispatch = crate::iteration::comprehension::strategies::for_name(strategy);
let Some(index_fn) = index_fn.filter(|idx| dispatch.accepts_input(Some(idx))) else {
return Err(RuntimeError::StrategyRejectsInput {
strategy,
index_fn: index_fn.cloned(),
});
};
Ok(dispatch.select_surviving(index_fn, cardinality, truncation, seed, survivors))
}
const SAMPLE_ROUNDS: u32 = 6;
const UNIT_SCALE: f64 = (1u64 << 53) as f64;
enum SampleAxis {
Discrete(Vec<RuntimeTuple>),
Continuous {
name: String,
interval: Interval,
measure: AxisMeasure,
},
}
#[derive(Default)]
struct SampleSpace {
axes: Vec<SampleAxis>,
predicates: Vec<String>,
}
impl SampleSpace {
fn realize(&self, mi: &[u64], strategy: StrategyName) -> RuntimeTuple {
let extrema = matches!(strategy, StrategyName::Extrema);
let discrete_count = self
.axes
.iter()
.filter(|a| matches!(a, SampleAxis::Discrete(_)))
.count();
let (mut d, mut c) = (0, if extrema { 0 } else { discrete_count });
let mut out = RuntimeTuple::new();
for axis in &self.axes {
match axis {
SampleAxis::Discrete(tuples) => {
let pos = mi.get(d).copied().unwrap_or(0) as usize;
d += 1;
if extrema {
c += 1;
}
if let Some(t) = tuples.get(pos) {
out.extend(t.iter().cloned());
}
}
SampleAxis::Continuous {
name,
interval,
measure,
} => {
let code = mi.get(c).copied().unwrap_or(0);
c += 1;
if extrema {
d += 1;
}
let x = if extrema {
measure.endpoint(interval, code == 1)
} else {
measure.map_unit(code as f64 / UNIT_SCALE, interval)
};
out.push((name.clone(), Value::F64(x)));
}
}
}
out
}
}
fn draw_sample(
index_fn: &IndexFn,
strategy: StrategyName,
count: Option<u64>,
seed: Option<u64>,
) -> Result<Vec<Vec<u64>>, RuntimeError> {
use crate::iteration::comprehension::strategies::{
extrema::extrema_multi_indices, halton::try_halton_multi_indices,
lhs::try_lhs_multi_indices, shuffle::try_shuffle_multi_indices,
sobol::try_sobol_multi_indices,
};
if matches!(strategy, StrategyName::Extrema) {
return Ok(extrema_multi_indices(index_fn, count));
}
let Some(n) = count else {
return Err(RuntimeError::OrderEval {
strategy,
message: "a continuous source has no finite tuple set; give the order a count, \
as in `order halton/16`"
.into(),
});
};
let drawn = match strategy {
StrategyName::Halton => try_halton_multi_indices(index_fn, Some(n)),
StrategyName::Sobol => try_sobol_multi_indices(index_fn, Some(n)),
StrategyName::Lhs => try_lhs_multi_indices(index_fn, Some(n), seed),
StrategyName::Shuffle => try_shuffle_multi_indices(index_fn, Some(n), seed),
other => {
return Err(RuntimeError::OrderEval {
strategy: other,
message: "a continuous source needs a sampling strategy: halton, sobol, lhs, \
shuffle, or extrema"
.into(),
});
}
};
drawn.map_err(|message| RuntimeError::OrderEval { strategy, message })
}
pub(crate) fn has_continuous_axis(c: &Comprehension) -> bool {
match c {
Comprehension::Clause { source, .. } => matches!(
source,
Source::ContinuousInterval { .. } | Source::Distribution { .. }
),
Comprehension::Cartesian { children } => children.iter().any(has_continuous_axis),
Comprehension::Filter { child, .. } => has_continuous_axis(child),
Comprehension::Zip { .. } | Comprehension::Union { .. } | Comprehension::Order { .. } => {
false
}
}
}
fn combine_cartesian_index_fn(children: &[Option<IndexFn>]) -> Option<IndexFn> {
let mut axis_sizes = Vec::new();
for opt in children {
match opt {
Some(IndexFn::Lattice { axis_sizes: a }) => axis_sizes.extend(a.iter().copied()),
Some(IndexFn::Lockstep { length }) => axis_sizes.push(*length),
_ => return None,
}
}
Some(IndexFn::Lattice { axis_sizes })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::iteration::comprehension::source::LiteralValue;
fn empty_kernel() -> Arc<PolydatKernel> {
Arc::new(crate::dsl::compile_polydat_interpreter("\n").unwrap())
}
fn canonical_with_k() -> Arc<PolydatKernel> {
Arc::new(crate::dsl::compile_polydat_interpreter("extern k: u64\n").unwrap())
}
fn clause(name: &str, source: Source) -> Comprehension {
Comprehension::Clause {
name: name.into(),
source,
}
}
fn empty_literal() -> Source {
Source::Literal { values: Vec::new() }
}
#[test]
fn every_leaf_reports_what_it_yielded() {
let comp = Comprehension::Cartesian {
children: vec![
clause(
"a",
Source::IntRange {
lo: 0,
hi: 3,
step: 1,
},
),
clause(
"b",
Source::Literal {
values: vec![LiteralValue::Int(7), LiteralValue::Int(8)],
},
),
],
};
let scope = empty_kernel();
let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
assert_eq!(out.tuples.len(), 6, "3 x 2");
assert_eq!(out.clauses.len(), 2, "one entry per leaf, in tree order");
assert_eq!(out.clauses[0].var, "a");
assert_eq!(out.clauses[0].values, 3);
assert_eq!(out.clauses[1].var, "b");
assert_eq!(out.clauses[1].evaluations, 3);
assert_eq!(out.clauses[1].values, 6);
}
#[test]
fn an_empty_clause_is_reached_and_yields_nothing() {
let comp = Comprehension::Cartesian {
children: vec![
clause(
"a",
Source::IntRange {
lo: 0,
hi: 2,
step: 1,
},
),
clause("b", empty_literal()),
],
};
let scope = empty_kernel();
let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
assert!(out.tuples.is_empty(), "an empty clause empties the product");
let culprits: Vec<&str> = out
.clauses
.iter()
.filter(|c| c.evaluations > 0 && c.values == 0)
.map(|c| c.var.as_str())
.collect();
assert_eq!(culprits, ["b"], "only the empty clause is named");
}
#[test]
fn a_clause_behind_an_empty_one_is_never_reached() {
let comp = Comprehension::Cartesian {
children: vec![
clause("outer", empty_literal()),
clause(
"inner",
Source::IntRange {
lo: 0,
hi: 9,
step: 1,
},
),
],
};
let scope = empty_kernel();
let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
assert!(out.tuples.is_empty());
let by = |v: &str| {
out.clauses
.iter()
.find(|c| c.var == v)
.expect("every leaf is present whether reached or not")
};
assert_eq!(by("outer").evaluations, 1);
assert_eq!(by("outer").values, 0);
assert_eq!(
by("inner").evaluations,
0,
"never reached: the cause is `outer`, not this"
);
assert_eq!(by("inner").values, 0);
}
#[test]
fn clauses_sharing_a_name_across_a_union_are_counted_apart() {
let comp = Comprehension::Union {
children: vec![
clause(
"k",
Source::Literal {
values: vec![LiteralValue::Int(1)],
},
),
clause("k", empty_literal()),
],
};
let scope = empty_kernel();
let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
assert_eq!(out.clauses.len(), 2, "two leaves, one name");
assert_eq!(out.clauses[0].values, 1);
assert_eq!(out.clauses[1].values, 0);
assert_eq!(out.clauses[1].evaluations, 1, "reached, and empty");
}
#[test]
fn the_plain_entry_point_agrees_with_the_reported_one() {
let comp = Comprehension::Cartesian {
children: vec![
clause(
"a",
Source::IntRange {
lo: 1,
hi: 4,
step: 1,
},
),
clause(
"b",
Source::Literal {
values: vec![LiteralValue::Int(5)],
},
),
],
};
let scope = empty_kernel();
let plain = evaluate_for_iteration(&comp, &*scope).unwrap();
let reported = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
assert_eq!(plain, reported.tuples);
}
#[test]
fn int_range_yields_values() {
let comp = Comprehension::Clause {
name: "k".into(),
source: Source::IntRange {
lo: 1,
hi: 5,
step: 1,
},
};
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 4);
assert_eq!(tuples[0][0].1, Value::U64(1));
assert_eq!(tuples[3][0].1, Value::U64(4));
}
#[test]
fn literal_list_yields_values() {
let comp = Comprehension::Clause {
name: "x".into(),
source: Source::Literal {
values: vec![LiteralValue::Int(10), LiteralValue::Int(20)],
},
};
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 2);
}
#[test]
fn cartesian_produces_product() {
let comp = Comprehension::cartesian(vec![
Comprehension::Clause {
name: "x".into(),
source: Source::IntRange {
lo: 1,
hi: 3,
step: 1,
},
},
Comprehension::Clause {
name: "y".into(),
source: Source::IntRange {
lo: 10,
hi: 30,
step: 10,
},
},
]);
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 4);
}
#[test]
fn union_produces_concatenation() {
let comp = Comprehension::union(vec![
Comprehension::Clause {
name: "k".into(),
source: Source::Literal {
values: vec![LiteralValue::Int(1)],
},
},
Comprehension::Clause {
name: "k".into(),
source: Source::Literal {
values: vec![LiteralValue::Int(10), LiteralValue::Int(20)],
},
},
]);
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 3);
}
#[test]
fn filter_drops_non_matching() {
let comp = Comprehension::filter(
Comprehension::Clause {
name: "k".into(),
source: Source::IntRange {
lo: 1,
hi: 6,
step: 1,
},
},
"{k} > 3",
);
let canonical = canonical_with_k();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 2);
}
#[test]
fn order_lex_truncate() {
let comp = Comprehension::order(
Comprehension::Clause {
name: "k".into(),
source: Source::IntRange {
lo: 1,
hi: 100,
step: 1,
},
},
StrategyName::Lex,
Some(5),
);
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 5);
}
#[test]
fn a_multi_name_head_is_one_lattice_axis() {
use crate::iteration::comprehension::strategy::ZipMode;
let lit = |name: &str, vs: &[i64]| Comprehension::Clause {
name: name.into(),
source: Source::Literal {
values: vs.iter().map(|v| LiteralValue::Int(*v)).collect(),
},
};
let comp = Comprehension::order(
Comprehension::cartesian(vec![
Comprehension::zip(vec![lit("a", &[1, 2]), lit("b", &[3, 4])], ZipMode::Strict),
lit("c", &[5, 6, 7]),
]),
StrategyName::Halton,
Some(6),
);
let tuples = evaluate_for_iteration(&comp, &*empty_kernel()).unwrap();
assert_eq!(
tuples.len(),
6,
"every tuple of the 2 x 3 product: {tuples:?}"
);
}
#[test]
fn extrema_over_cartesian_uses_indexed_form() {
let comp = Comprehension::order(
Comprehension::cartesian(vec![
Comprehension::Clause {
name: "k".into(),
source: Source::Literal {
values: vec![
LiteralValue::Int(1),
LiteralValue::Int(2),
LiteralValue::Int(3),
],
},
},
Comprehension::Clause {
name: "limit".into(),
source: Source::Literal {
values: vec![
LiteralValue::Int(10),
LiteralValue::Int(20),
LiteralValue::Int(30),
],
},
},
]),
StrategyName::Extrema,
Some(1),
);
let canonical = empty_kernel();
let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
assert_eq!(tuples.len(), 4);
for t in &tuples {
assert_eq!(t.len(), 2);
let k = match &t[0].1 {
Value::U64(n) => *n,
other => panic!("expected u64 k, got {other:?}"),
};
let lim = match &t[1].1 {
Value::U64(n) => *n,
other => panic!("expected u64 limit, got {other:?}"),
};
assert!(k == 1 || k == 3, "expected extreme k, got {k}");
assert!(lim == 10 || lim == 30, "expected extreme limit, got {lim}");
}
}
}