1use std::sync::Arc;
2
3use laddu_compile::CompiledQuery;
4use laddu_data::{
5 LadduDataError, LadduDataResult,
6 data::{Dataset, EventBatch},
7 io::{EventBatchIter, EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
8 schema::Schema,
9};
10use laddu_expr::{Expr, ValueKind};
11use laddu_physics::{
12 binning::{BinningAxis, FinalUpperEdge},
13 histogram::Histogram,
14 joint_histogram::JointHistogram,
15};
16use num::complex::Complex64;
17use serde::{Deserialize, Deserializer, Serialize};
18
19use crate::{Execution, PreparedModel, RuntimeError, RuntimeResult};
20
21#[derive(Copy, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
23pub enum Comparison {
24 Lt,
26 Le,
28 Gt,
30 Ge,
32 Eq,
34 Ne,
36}
37
38#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
40pub enum IntervalClosure {
41 Open,
43 LeftClosed,
45 RightClosed,
47 #[default]
49 Closed,
50}
51
52#[derive(Clone, Debug)]
54pub enum Predicate {
55 Compare {
57 lhs: Expr,
59 op: Comparison,
61 rhs: Expr,
63 },
64 And(Box<Self>, Box<Self>),
66 Or(Box<Self>, Box<Self>),
68 Not(Box<Self>),
70 Between {
72 value: Expr,
74 lower: Expr,
76 upper: Expr,
78 closure: IntervalClosure,
80 },
81}
82
83impl Predicate {
84 pub fn compare(lhs: impl Into<Expr>, op: Comparison, rhs: impl Into<Expr>) -> Self {
86 Self::Compare {
87 lhs: lhs.into(),
88 op,
89 rhs: rhs.into(),
90 }
91 }
92
93 pub fn lt(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
95 Self::compare(lhs, Comparison::Lt, rhs)
96 }
97 pub fn le(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
99 Self::compare(lhs, Comparison::Le, rhs)
100 }
101 pub fn gt(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
103 Self::compare(lhs, Comparison::Gt, rhs)
104 }
105 pub fn ge(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
107 Self::compare(lhs, Comparison::Ge, rhs)
108 }
109 pub fn eq(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
111 Self::compare(lhs, Comparison::Eq, rhs)
112 }
113 pub fn ne(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Self {
115 Self::compare(lhs, Comparison::Ne, rhs)
116 }
117 pub fn and(self, rhs: Self) -> Self {
119 Self::And(Box::new(self), Box::new(rhs))
120 }
121 pub fn or(self, rhs: Self) -> Self {
123 Self::Or(Box::new(self), Box::new(rhs))
124 }
125 pub fn between(value: impl Into<Expr>, lower: impl Into<Expr>, upper: impl Into<Expr>) -> Self {
127 Self::between_with(value, lower, upper, IntervalClosure::Closed)
128 }
129 pub fn between_with(
131 value: impl Into<Expr>,
132 lower: impl Into<Expr>,
133 upper: impl Into<Expr>,
134 closure: IntervalClosure,
135 ) -> Self {
136 Self::Between {
137 value: value.into(),
138 lower: lower.into(),
139 upper: upper.into(),
140 closure,
141 }
142 }
143}
144
145impl std::ops::Not for Predicate {
146 type Output = Self;
147 fn not(self) -> Self::Output {
148 Self::Not(Box::new(self))
149 }
150}
151
152#[derive(Clone, Debug, PartialEq, Serialize)]
154#[serde(transparent)]
155pub struct BinSpec {
156 axis: BinningAxis,
157}
158
159impl<'de> Deserialize<'de> for BinSpec {
160 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
161 where
162 D: Deserializer<'de>,
163 {
164 let edges = Vec::<f64>::deserialize(deserializer)?;
165 Self::edges(edges).map_err(serde::de::Error::custom)
166 }
167}
168
169impl BinSpec {
170 pub fn uniform(count: usize, min: f64, max: f64) -> RuntimeResult<Self> {
177 Ok(Self {
178 axis: BinningAxis::uniform(count, min, max)
179 .map_err(|error| query_error(error.to_string()))?,
180 })
181 }
182
183 pub fn edges(edges: impl IntoIterator<Item = f64>) -> RuntimeResult<Self> {
190 Ok(Self {
191 axis: BinningAxis::new(edges).map_err(|error| query_error(error.to_string()))?,
192 })
193 }
194
195 pub fn bin_count(&self) -> usize {
197 self.axis.bin_count()
198 }
199 pub fn edges_slice(&self) -> &[f64] {
201 self.axis.edges()
202 }
203
204 fn index(&self, value: f64) -> Option<usize> {
205 self.axis.index(value, FinalUpperEdge::Inclusive)
206 }
207}
208
209#[derive(Clone)]
211pub struct DatasetBin {
212 index: usize,
213 lower: f64,
214 upper: f64,
215 dataset: Dataset,
216}
217
218impl DatasetBin {
219 pub fn index(&self) -> usize {
221 self.index
222 }
223 pub fn lower(&self) -> f64 {
225 self.lower
226 }
227 pub fn upper(&self) -> f64 {
229 self.upper
230 }
231 pub fn dataset(&self) -> &Dataset {
233 &self.dataset
234 }
235 pub fn into_dataset(self) -> Dataset {
237 self.dataset
238 }
239}
240
241pub trait DatasetExprExt {
243 fn validate_real_expressions(
249 &self,
250 expressions: &[Expr],
251 execution: &Execution,
252 ) -> RuntimeResult<()>;
253 fn evaluate_expr(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<Complex64>>;
260 fn evaluate_real(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<f64>>;
267 fn visit_real_chunks(
275 &self,
276 expressions: &[Expr],
277 execution: &Execution,
278 chunk_size: usize,
279 consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
280 ) -> RuntimeResult<()>;
281 fn histogram(
292 &self,
293 expr: &Expr,
294 bins: BinSpec,
295 event_weights: bool,
296 weight: Option<&Expr>,
297 execution: &Execution,
298 ) -> RuntimeResult<Histogram>;
299 fn joint_histogram(
306 &self,
307 axes: &[Expr],
308 bins: Vec<BinSpec>,
309 event_weights: bool,
310 weight: Option<&Expr>,
311 execution: &Execution,
312 ) -> RuntimeResult<JointHistogram>;
313 fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset>;
320 fn bin_by(
327 &self,
328 expr: &Expr,
329 bins: BinSpec,
330 execution: &Execution,
331 ) -> RuntimeResult<Vec<DatasetBin>>;
332}
333
334impl DatasetExprExt for Dataset {
335 fn validate_real_expressions(
336 &self,
337 expressions: &[Expr],
338 execution: &Execution,
339 ) -> RuntimeResult<()> {
340 if expressions.is_empty() {
341 return Err(query_error(
342 "real expression validation needs at least one expression",
343 ));
344 }
345 QueryExprSet::prepare(expressions.to_vec(), execution, true)?;
346 Ok(())
347 }
348
349 fn evaluate_expr(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<Complex64>> {
350 let query = QueryExprSet::prepare(vec![expr.clone()], execution, false)?;
351 let mut output = Vec::new();
352 for batch in self.batches().map_err(data_error)? {
353 output.extend(
354 query.evaluate_batch(&batch.map_err(data_error)?)?[0]
355 .iter()
356 .copied(),
357 );
358 }
359 Ok(output)
360 }
361
362 fn evaluate_real(&self, expr: &Expr, execution: &Execution) -> RuntimeResult<Vec<f64>> {
363 let query = QueryExprSet::prepare(vec![expr.clone()], execution, true)?;
364 let mut output = Vec::new();
365 for batch in self.batches().map_err(data_error)? {
366 output.extend(
367 query.evaluate_batch(&batch.map_err(data_error)?)?[0]
368 .iter()
369 .copied()
370 .map(|v| v.re),
371 );
372 }
373 Ok(output)
374 }
375
376 fn visit_real_chunks(
377 &self,
378 expressions: &[Expr],
379 execution: &Execution,
380 chunk_size: usize,
381 mut consume: impl FnMut(usize, &[Vec<f64>]) -> RuntimeResult<()>,
382 ) -> RuntimeResult<()> {
383 if chunk_size == 0 || expressions.is_empty() {
384 return Err(query_error(
385 "real chunk evaluation needs expressions and positive chunk size",
386 ));
387 }
388 let query = QueryExprSet::prepare(expressions.to_vec(), execution, true)?;
389 let mut offset = 0;
390 for batch in self.batches().map_err(data_error)? {
391 let batch = batch.map_err(data_error)?;
392 for start in (0..batch.len()).step_by(chunk_size) {
393 let end = (start + chunk_size).min(batch.len());
394 let values = query
395 .evaluate_batch(&batch.slice(start, end))?
396 .into_iter()
397 .map(|column| column.into_iter().map(|value| value.re).collect::<Vec<_>>())
398 .collect::<Vec<_>>();
399 consume(offset, &values)?;
400 offset += end - start;
401 }
402 }
403 Ok(())
404 }
405
406 fn histogram(
407 &self,
408 expr: &Expr,
409 bins: BinSpec,
410 event_weights: bool,
411 weight: Option<&Expr>,
412 execution: &Execution,
413 ) -> RuntimeResult<Histogram> {
414 let mut expressions = vec![expr.clone()];
415 expressions.extend(weight.cloned());
416 let query = QueryExprSet::prepare(expressions, execution, true)?;
417 let mut histogram = Histogram::empty_with_edges(bins.edges_slice().to_vec())
418 .map_err(|error| query_error(error.to_string()))?;
419
420 for batch in self.batches().map_err(data_error)? {
421 let batch = batch.map_err(data_error)?;
422 let values = query.evaluate_batch(&batch)?;
423 for row in 0..batch.len() {
424 let base_weight = if event_weights {
425 batch.weights_at(row)
426 } else {
427 1.0
428 };
429 let custom_weight = values.get(1).map_or(1.0, |weights| weights[row].re);
430 histogram
431 .fill_weighted(values[0][row].re, base_weight * custom_weight)
432 .map_err(|error| query_error(error.to_string()))?;
433 }
434 }
435 Ok(histogram)
436 }
437
438 fn joint_histogram(
439 &self,
440 axes: &[Expr],
441 bins: Vec<BinSpec>,
442 event_weights: bool,
443 weight: Option<&Expr>,
444 execution: &Execution,
445 ) -> RuntimeResult<JointHistogram> {
446 if axes.is_empty() || axes.len() != bins.len() {
447 return Err(query_error(
448 "joint histogram requires one bin specification per non-empty ordered axis",
449 ));
450 }
451 let edge_vectors = bins
452 .iter()
453 .map(|bins| bins.edges_slice().to_vec())
454 .collect();
455 let mut histogram =
456 JointHistogram::empty(edge_vectors).map_err(|error| query_error(error.to_string()))?;
457 let mut expressions = axes.to_vec();
458 expressions.extend(weight.cloned());
459 let query = QueryExprSet::prepare(expressions, execution, true)?;
460 let mut coordinates = vec![0.0; axes.len()];
461 for batch in self.batches().map_err(data_error)? {
462 let batch = batch.map_err(data_error)?;
463 let values = query.evaluate_batch(&batch)?;
464 for row in 0..batch.len() {
465 for (axis, coordinate) in coordinates.iter_mut().enumerate() {
466 *coordinate = values[axis][row].re;
467 }
468 let base_weight = if event_weights {
469 batch.weights_at(row)
470 } else {
471 1.0
472 };
473 let custom_weight = values
474 .get(axes.len())
475 .map_or(1.0, |weights| weights[row].re);
476 histogram
477 .fill_weighted(&coordinates, base_weight * custom_weight)
478 .map_err(|error| query_error(error.to_string()))?;
479 }
480 }
481 Ok(histogram)
482 }
483
484 fn select(&self, predicate: &Predicate, execution: &Execution) -> RuntimeResult<Dataset> {
485 let compiled = CompiledPredicate::prepare(predicate, execution)?;
486 Ok(self.with_derived_source(QuerySource {
487 source: self.clone(),
488 filter: QueryFilter::Predicate(Arc::new(compiled)),
489 }))
490 }
491
492 fn bin_by(
493 &self,
494 expr: &Expr,
495 bins: BinSpec,
496 execution: &Execution,
497 ) -> RuntimeResult<Vec<DatasetBin>> {
498 let query = QueryExprSet::prepare(vec![expr.clone()], execution, true)?;
499 let schema = self.schema().map_err(data_error)?;
500 let mut partitions = vec![Vec::new(); bins.bin_count()];
501 for batch in self.batches().map_err(data_error)? {
502 let batch = batch.map_err(data_error)?;
503 let mut rows = vec![Vec::new(); bins.bin_count()];
504 for (row, value) in query.evaluate_batch(&batch)?[0].iter().copied().enumerate() {
505 if let Some(index) = bins.index(value.re) {
506 rows[index].push(row);
507 }
508 }
509 for (partition, rows) in partitions.iter_mut().zip(rows) {
510 if !rows.is_empty() {
511 partition.push(batch.select(&rows));
512 }
513 }
514 }
515
516 partitions
517 .into_iter()
518 .enumerate()
519 .map(|(index, batches)| {
520 let source = if batches.is_empty() {
521 MemorySource::empty(Arc::clone(&schema))
522 } else {
523 MemorySource::from_batches(batches).map_err(data_error)?
524 };
525 Ok(DatasetBin {
526 index,
527 lower: bins.edges_slice()[index],
528 upper: bins.edges_slice()[index + 1],
529 dataset: self.with_derived_source(source),
530 })
531 })
532 .collect()
533 }
534}
535
536struct QueryExpr {
537 model: PreparedModel,
538 params: laddu_expr::parameters::ParamValues,
539 outputs: Vec<laddu_expr::ExprId>,
540}
541
542struct QueryExprSet {
543 shared: QueryExprStorage,
544 outputs: usize,
545}
546
547enum QueryExprStorage {
548 Shared(QueryExpr),
549 Separate(Vec<QueryExpr>),
550}
551
552impl QueryExprSet {
553 fn prepare(
554 expressions: Vec<Expr>,
555 execution: &Execution,
556 require_real: bool,
557 ) -> RuntimeResult<Self> {
558 let expression_count = expressions.len();
559 let compiled = CompiledQuery::from_exprs(expressions.clone())
560 .map_err(|error| query_error(error.to_string()))?;
561 let model = compiled.model();
562 let outputs = compiled.outputs();
563 if outputs.len() != expression_count {
564 return Err(query_error(
565 "compiled query output count changed during lowering",
566 ));
567 }
568 for element in outputs {
569 let value_kind = model
570 .node_facts(*element)
571 .map(|facts| facts.value_kind)
572 .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
573 if value_kind != ValueKind::Real
574 && (require_real || !matches!(value_kind, ValueKind::Complex))
575 {
576 return Err(query_error(if require_real {
577 "this dataset operation requires a real-valued expression"
578 } else {
579 "dataset expressions must be scalar"
580 }));
581 }
582 }
583 if model.params().n_free() != 0 {
584 return Err(query_error(
585 "dataset expressions cannot contain free parameters",
586 ));
587 }
588 let params = model.params().default_values();
589 let plan = match PreparedModel::prepare(model, execution) {
590 Ok(plan) => QueryExprStorage::Shared(QueryExpr {
591 model: plan,
592 params,
593 outputs: outputs.to_vec(),
594 }),
595 Err(shared_error) => {
596 if !may_fallback_to_scalar(execution, &shared_error) {
597 return Err(shared_error);
598 }
599 let separate = expressions
600 .iter()
601 .map(|expr| QueryExpr::prepare(expr, execution, require_real))
602 .collect::<RuntimeResult<Vec<_>>>();
603 QueryExprStorage::Separate(separate?)
604 }
605 };
606 Ok(Self {
607 shared: plan,
608 outputs: expression_count,
609 })
610 }
611
612 fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Vec<Complex64>>> {
613 let values = match &self.shared {
614 QueryExprStorage::Shared(query) => {
615 query
616 .model
617 .evaluate_batch_outputs(&query.params, batch, &query.outputs)?
618 }
619 QueryExprStorage::Separate(queries) => queries
620 .iter()
621 .map(|query| query.evaluate_batch(batch))
622 .collect::<RuntimeResult<Vec<_>>>()?,
623 };
624 if values.len() != self.outputs {
625 return Err(query_error(
626 "compiled query returned an unexpected output count",
627 ));
628 }
629 Ok(values)
630 }
631}
632
633impl QueryExpr {
634 fn prepare(expr: &Expr, execution: &Execution, require_real: bool) -> RuntimeResult<Self> {
635 if expr.shape().map_err(|e| query_error(e.to_string()))? != laddu_expr::ExprShape::Scalar {
636 return Err(query_error("dataset expressions must be scalar"));
637 }
638 let compiled = laddu_compile::CompiledModel::from_expr(expr)
639 .map_err(|error| query_error(error.to_string()))?;
640 let value_kind = compiled
641 .node_facts(compiled.graph().root())
642 .map(|facts| facts.value_kind)
643 .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
644 if require_real && value_kind == ValueKind::Complex {
645 return Err(query_error(
646 "this dataset operation requires a real-valued expression",
647 ));
648 }
649 if compiled.params().n_free() != 0 {
650 return Err(query_error(
651 "dataset expressions cannot contain free parameters",
652 ));
653 }
654 let params = compiled.params().default_values();
655 let model = PreparedModel::prepare(&compiled, execution)?;
656 Ok(Self {
657 model,
658 params,
659 outputs: Vec::new(),
660 })
661 }
662
663 fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Complex64>> {
664 self.model.evaluate_batch(&self.params, batch)
665 }
666}
667
668struct CompiledPredicate {
669 expressions: QueryExprSet,
670 program: PredicateProgram,
671}
672
673enum PredicateProgram {
674 Compare {
675 lhs: usize,
676 op: Comparison,
677 rhs: usize,
678 },
679 And(Box<Self>, Box<Self>),
680 Or(Box<Self>, Box<Self>),
681 Not(Box<Self>),
682 Between {
683 value: usize,
684 lower: usize,
685 upper: usize,
686 closure: IntervalClosure,
687 },
688}
689
690impl CompiledPredicate {
691 fn prepare(predicate: &Predicate, execution: &Execution) -> RuntimeResult<Self> {
692 let mut expressions = Vec::new();
693 let program = Self::compile_program(predicate, &mut expressions);
694 Ok(Self {
695 expressions: QueryExprSet::prepare(expressions, execution, true)?,
696 program,
697 })
698 }
699
700 fn compile_program(predicate: &Predicate, expressions: &mut Vec<Expr>) -> PredicateProgram {
701 let leaf = |expr: &Expr, expressions: &mut Vec<Expr>| {
702 let index = expressions.len();
703 expressions.push(expr.clone());
704 index
705 };
706 match predicate {
707 Predicate::Compare { lhs, op, rhs } => PredicateProgram::Compare {
708 lhs: leaf(lhs, expressions),
709 op: *op,
710 rhs: leaf(rhs, expressions),
711 },
712 Predicate::And(lhs, rhs) => PredicateProgram::And(
713 Box::new(Self::compile_program(lhs, expressions)),
714 Box::new(Self::compile_program(rhs, expressions)),
715 ),
716 Predicate::Or(lhs, rhs) => PredicateProgram::Or(
717 Box::new(Self::compile_program(lhs, expressions)),
718 Box::new(Self::compile_program(rhs, expressions)),
719 ),
720 Predicate::Not(inner) => {
721 PredicateProgram::Not(Box::new(Self::compile_program(inner, expressions)))
722 }
723 Predicate::Between {
724 value,
725 lower,
726 upper,
727 closure,
728 } => PredicateProgram::Between {
729 value: leaf(value, expressions),
730 lower: leaf(lower, expressions),
731 upper: leaf(upper, expressions),
732 closure: *closure,
733 },
734 }
735 }
736
737 fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
738 let values = self.expressions.evaluate_batch(batch)?;
739 Ok((0..batch.len())
740 .filter(|row| Self::evaluate_row(&self.program, &values, *row))
741 .collect())
742 }
743
744 fn evaluate_row(program: &PredicateProgram, values: &[Vec<Complex64>], row: usize) -> bool {
745 match program {
746 PredicateProgram::Compare { lhs, op, rhs } => {
747 compare(values[*lhs][row].re, *op, values[*rhs][row].re)
748 }
749 PredicateProgram::And(lhs, rhs) => {
750 Self::evaluate_row(lhs, values, row) && Self::evaluate_row(rhs, values, row)
751 }
752 PredicateProgram::Or(lhs, rhs) => {
753 Self::evaluate_row(lhs, values, row) || Self::evaluate_row(rhs, values, row)
754 }
755 PredicateProgram::Not(inner) => !Self::evaluate_row(inner, values, row),
756 PredicateProgram::Between {
757 value,
758 lower,
759 upper,
760 closure,
761 } => {
762 let lower_op = match closure {
763 IntervalClosure::Open | IntervalClosure::RightClosed => Comparison::Gt,
764 IntervalClosure::LeftClosed | IntervalClosure::Closed => Comparison::Ge,
765 };
766 let upper_op = match closure {
767 IntervalClosure::Open | IntervalClosure::LeftClosed => Comparison::Lt,
768 IntervalClosure::RightClosed | IntervalClosure::Closed => Comparison::Le,
769 };
770 compare(values[*value][row].re, lower_op, values[*lower][row].re)
771 && compare(values[*value][row].re, upper_op, values[*upper][row].re)
772 }
773 }
774 }
775}
776
777fn compare(lhs: f64, op: Comparison, rhs: f64) -> bool {
778 if lhs.is_nan() || rhs.is_nan() {
779 return false;
780 }
781 match op {
782 Comparison::Lt => lhs < rhs,
783 Comparison::Le => lhs <= rhs,
784 Comparison::Gt => lhs > rhs,
785 Comparison::Ge => lhs >= rhs,
786 Comparison::Eq => lhs == rhs,
787 Comparison::Ne => lhs != rhs,
788 }
789}
790
791#[derive(Clone)]
792struct QuerySource {
793 source: Dataset,
794 filter: QueryFilter,
795}
796
797#[derive(Clone)]
798enum QueryFilter {
799 Predicate(Arc<CompiledPredicate>),
800}
801
802impl EventSource for QuerySource {
803 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
804 self.source.schema()
805 }
806
807 fn capabilities(&self) -> SourceCapabilities {
808 let source = self.source.capabilities();
809 SourceCapabilities {
810 exact_len: false,
811 exact_weighted_total: false,
812 random_access: false,
813 deterministic_partitioning: source.deterministic_partitioning,
814 predicate_pushdown: false,
815 projection_pushdown: false,
816 streaming: true,
817 }
818 }
819
820 fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
821 let batches = self.source.stream_with_plan(plan)?;
822 let filter = self.filter.clone();
823 Ok(Box::new(batches.filter_map(move |batch| {
824 let batch = match batch {
825 Ok(batch) => batch,
826 Err(error) => return Some(Err(error)),
827 };
828 let rows = match filter.rows(&batch) {
829 Ok(rows) => rows,
830 Err(error) => return Some(Err(LadduDataError::Source(error.to_string()))),
831 };
832 (!rows.is_empty()).then(|| Ok(batch.select(&rows)))
833 })))
834 }
835}
836
837impl QueryFilter {
838 fn rows(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
839 match self {
840 Self::Predicate(predicate) => predicate.evaluate_batch(batch),
841 }
842 }
843}
844
845fn query_error(message: impl Into<String>) -> RuntimeError {
846 RuntimeError::InvalidShape {
847 index: 0,
848 message: message.into(),
849 }
850}
851
852fn may_fallback_to_scalar(execution: &Execution, error: &RuntimeError) -> bool {
853 let cpu_f32 = matches!(
854 error,
855 RuntimeError::Execution(crate::ExecutionError::UnsupportedCpuF32Model)
856 );
857 #[cfg(feature = "wgpu")]
858 {
859 cpu_f32 || (execution.wgpu_context().is_some() && matches!(error, RuntimeError::Wgpu(_)))
860 }
861 #[cfg(not(feature = "wgpu"))]
862 {
863 let _ = execution;
864 cpu_f32
865 }
866}
867
868fn data_error(error: impl ToString) -> RuntimeError {
869 RuntimeError::Data(error.to_string())
870}
871
872#[cfg(test)]
873mod tests {
874 use super::*;
875 use crate::{CpuOptions, Device, ExecutionOptions, Precision};
876 #[cfg(feature = "jit")]
877 use crate::{JitPolicy, ThreadPolicy};
878 use laddu_compile::CompiledModel;
879 use laddu_data::{
880 data::{EventBatch, OwnedEvent},
881 io::{EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
882 schema::Schema,
883 };
884 use laddu_expr::{complex, event_scalar};
885 use std::sync::atomic::{AtomicUsize, Ordering};
886
887 #[test]
888 fn bin_spec_roundtrip_preserves_validation() {
889 let bins = BinSpec::edges([-1.0, 0.0, 2.0]).unwrap();
890 let json = serde_json::to_string(&bins).unwrap();
891 assert_eq!(serde_json::from_str::<BinSpec>(&json).unwrap(), bins);
892 assert!(serde_json::from_str::<BinSpec>("[0.0,0.0]").is_err());
893 }
894
895 #[derive(Clone)]
896 struct CountingSource {
897 inner: MemorySource,
898 reads: Arc<AtomicUsize>,
899 }
900
901 impl EventSource for CountingSource {
902 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
903 EventSource::schema(&self.inner)
904 }
905
906 fn capabilities(&self) -> SourceCapabilities {
907 self.inner.capabilities()
908 }
909
910 fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
911 self.reads.fetch_add(1, Ordering::Relaxed);
912 self.inner.batches(plan)
913 }
914 }
915
916 #[derive(Clone)]
917 struct FailingSource {
918 schema: Arc<Schema>,
919 }
920
921 impl EventSource for FailingSource {
922 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
923 Ok(Arc::clone(&self.schema))
924 }
925
926 fn capabilities(&self) -> SourceCapabilities {
927 SourceCapabilities {
928 exact_len: false,
929 exact_weighted_total: false,
930 random_access: false,
931 deterministic_partitioning: true,
932 predicate_pushdown: false,
933 projection_pushdown: false,
934 streaming: true,
935 }
936 }
937
938 fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
939 Ok(Box::new(std::iter::once(Err(LadduDataError::Source(
940 "query source failed".into(),
941 )))))
942 }
943 }
944
945 fn capability_tuple(
946 capabilities: SourceCapabilities,
947 ) -> (bool, bool, bool, bool, bool, bool, bool) {
948 (
949 capabilities.exact_len,
950 capabilities.exact_weighted_total,
951 capabilities.random_access,
952 capabilities.deterministic_partitioning,
953 capabilities.predicate_pushdown,
954 capabilities.projection_pushdown,
955 capabilities.streaming,
956 )
957 }
958
959 fn dataset() -> Dataset {
960 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
961 Dataset::from_events(
962 schema,
963 [
964 OwnedEvent::weighted(vec![], vec![-1.0], 0.5),
965 OwnedEvent::weighted(vec![], vec![0.0], 1.0),
966 OwnedEvent::weighted(vec![], vec![1.0], 1.5),
967 OwnedEvent::weighted(vec![], vec![2.0], 2.0),
968 ],
969 )
970 .unwrap()
971 }
972
973 #[test]
974 fn evaluates_selects_and_bins_dataset_expressions() {
975 let dataset = dataset().chunked(1).unwrap();
976 let execution = Execution::default();
977 let x = event_scalar("x");
978 assert_eq!(
979 dataset.evaluate_real(&x, &execution).unwrap(),
980 vec![-1.0, 0.0, 1.0, 2.0]
981 );
982
983 let selected = dataset
984 .select(
985 &Predicate::ge(x.clone(), 0.0).and(Predicate::lt(x.clone(), 2.0)),
986 &execution,
987 )
988 .unwrap();
989 assert_eq!(
990 selected.map_events(|event| event.scalar(0)).unwrap(),
991 vec![0.0, 1.0]
992 );
993 assert_eq!(selected.sum_weights().unwrap(), 2.5);
994
995 let bins = dataset
996 .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
997 .unwrap();
998 assert_eq!(bins.len(), 2);
999 assert_eq!(
1000 bins[0]
1001 .dataset()
1002 .map_events(|event| event.scalar(0))
1003 .unwrap(),
1004 vec![0.0]
1005 );
1006 assert_eq!(
1007 bins[1]
1008 .dataset()
1009 .map_events(|event| event.scalar(0))
1010 .unwrap(),
1011 vec![1.0, 2.0]
1012 );
1013 }
1014
1015 #[test]
1016 fn dataset_histogram_uses_event_weights_and_excludes_the_final_upper_edge() {
1017 let histogram = dataset()
1018 .histogram(
1019 &event_scalar("x"),
1020 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1021 true,
1022 None,
1023 &Execution::default(),
1024 )
1025 .unwrap();
1026
1027 assert_eq!(histogram.counts(), [1.0, 1.5]);
1028 assert_eq!(histogram.sum_squared_weights(), [1.0, 2.25]);
1029 assert_eq!(histogram.underflow(), 0.5);
1030 assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.25));
1031 assert_eq!(histogram.overflow(), 2.0);
1032 assert_eq!(histogram.overflow_sum_squared_weights(), Some(4.0));
1033 }
1034
1035 #[test]
1036 fn dataset_joint_histogram_uses_row_major_bins_and_aggregates_invalid_events() {
1037 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x", "y"], true).unwrap());
1038 let dataset = Dataset::from_events(
1039 schema,
1040 [
1041 OwnedEvent::weighted(vec![], vec![0.0, 10.0], 1.0),
1042 OwnedEvent::weighted(vec![], vec![1.0, 10.0], -2.0),
1043 OwnedEvent::weighted(vec![], vec![0.0, 20.0], 3.0),
1044 OwnedEvent::weighted(vec![], vec![f64::NAN, 10.0], 4.0),
1045 ],
1046 )
1047 .unwrap();
1048
1049 let histogram = dataset
1050 .joint_histogram(
1051 &[event_scalar("x"), event_scalar("y")],
1052 vec![
1053 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1054 BinSpec::edges([0.0, 15.0, 25.0]).unwrap(),
1055 ],
1056 true,
1057 None,
1058 &Execution::default(),
1059 )
1060 .unwrap();
1061
1062 assert_eq!(histogram.shape(), [2, 2]);
1063 assert_eq!(histogram.values(), [1.0, 3.0, -2.0, 0.0]);
1064 assert_eq!(histogram.sum_squared_weights(), [1.0, 9.0, 4.0, 0.0]);
1065 assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
1066 assert_eq!(histogram.diagnostics().out_of_range_count(), 0);
1067 }
1068
1069 #[test]
1070 fn dataset_histogram_multiplies_the_selected_base_and_custom_weights() {
1071 let x = event_scalar("x");
1072 let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1073 let execution = Execution::default();
1074
1075 let weighted = dataset()
1076 .histogram(&x, bins.clone(), true, Some(&(x.clone() + 2.0)), &execution)
1077 .unwrap();
1078 assert_eq!(weighted.counts(), [0.5, 6.5]);
1079 assert_eq!(weighted.sum_squared_weights(), [0.25, 24.25]);
1080 assert_eq!(weighted.overflow(), 8.0);
1081
1082 let custom_only = dataset()
1083 .histogram(&x, bins, false, Some(&(x.clone() + 2.0)), &execution)
1084 .unwrap();
1085 assert_eq!(custom_only.counts(), [1.0, 5.0]);
1086 assert_eq!(custom_only.sum_squared_weights(), [1.0, 13.0]);
1087 assert_eq!(custom_only.overflow(), 4.0);
1088 }
1089
1090 #[test]
1091 fn dataset_histogram_preserves_view_source_and_memory_semantics() {
1092 let source = dataset();
1093 let batch = source.batches().unwrap().next().unwrap().unwrap();
1094 let reads = Arc::new(AtomicUsize::new(0));
1095 let counted = Dataset::new(CountingSource {
1096 inner: MemorySource::new(batch),
1097 reads: Arc::clone(&reads),
1098 })
1099 .streaming()
1100 .chunked(1)
1101 .unwrap();
1102 let selected = counted
1103 .select(
1104 &Predicate::ge(event_scalar("x"), 0.0),
1105 &Execution::default(),
1106 )
1107 .unwrap();
1108
1109 let histogram = selected
1110 .histogram(
1111 &event_scalar("x"),
1112 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1113 true,
1114 None,
1115 &Execution::default(),
1116 )
1117 .unwrap();
1118
1119 assert_eq!(histogram.counts(), [1.0, 1.5]);
1120 assert_eq!(histogram.overflow(), 2.0);
1121 assert_eq!(reads.load(Ordering::Relaxed), 1);
1122 }
1123
1124 #[test]
1125 fn dataset_histogram_matches_across_memory_policies_and_chunking() {
1126 let source = dataset();
1127 let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1128 let execution = Execution::default();
1129 let expected = source
1130 .clone()
1131 .resident()
1132 .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1133 .unwrap();
1134
1135 for candidate in [
1136 source.clone().streaming(),
1137 source.clone().resident().chunked(1).unwrap(),
1138 source.streaming().chunked(2).unwrap(),
1139 ] {
1140 let actual = candidate
1141 .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1142 .unwrap();
1143 assert_eq!(actual, expected);
1144 }
1145 }
1146
1147 #[cfg(feature = "jit")]
1148 #[test]
1149 fn dataset_histogram_matches_cpu_interpreter_and_jit_backends() {
1150 let source = dataset();
1151 let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1152 let execution = |jit| {
1153 Execution::local(ExecutionOptions {
1154 device: Device::Cpu(CpuOptions {
1155 threads: ThreadPolicy::Serial,
1156 jit,
1157 }),
1158 precision: Precision::F64,
1159 ..ExecutionOptions::default()
1160 })
1161 .unwrap()
1162 };
1163 let interpreted = source
1164 .histogram(
1165 &event_scalar("x"),
1166 bins.clone(),
1167 true,
1168 None,
1169 &execution(JitPolicy::Disabled),
1170 )
1171 .unwrap();
1172 let compiled = source
1173 .histogram(
1174 &event_scalar("x"),
1175 bins,
1176 true,
1177 None,
1178 &execution(JitPolicy::Enabled),
1179 )
1180 .unwrap();
1181
1182 assert_eq!(compiled, interpreted);
1183 }
1184
1185 #[test]
1186 fn dataset_histogram_reports_invalid_values_and_source_failures() {
1187 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1188 let nonfinite = Dataset::from_events(
1189 Arc::clone(&schema),
1190 [OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0)],
1191 )
1192 .unwrap();
1193 let error = nonfinite
1194 .histogram(
1195 &event_scalar("x"),
1196 BinSpec::edges([0.0, 1.0]).unwrap(),
1197 true,
1198 None,
1199 &Execution::default(),
1200 )
1201 .unwrap_err();
1202 assert!(
1203 error.to_string().contains("expected finite, got NaN"),
1204 "unexpected error: {error}",
1205 );
1206
1207 let source_error = Dataset::new(FailingSource { schema })
1208 .histogram(
1209 &event_scalar("x"),
1210 BinSpec::edges([0.0, 1.0]).unwrap(),
1211 true,
1212 None,
1213 &Execution::default(),
1214 )
1215 .unwrap_err();
1216 assert!(source_error.to_string().contains("query source failed"));
1217 }
1218
1219 #[test]
1220 fn empty_dataset_histogram_is_a_valid_empirical_histogram() {
1221 let histogram = dataset()
1222 .empty_derived()
1223 .unwrap()
1224 .histogram(
1225 &event_scalar("x"),
1226 BinSpec::edges([0.0, 1.0, 3.0]).unwrap(),
1227 true,
1228 None,
1229 &Execution::default(),
1230 )
1231 .unwrap();
1232
1233 assert_eq!(histogram.counts(), [0.0, 0.0]);
1234 assert_eq!(histogram.sum_squared_weights(), [0.0, 0.0]);
1235 assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.0));
1236 assert_eq!(histogram.overflow_sum_squared_weights(), Some(0.0));
1237 }
1238
1239 #[test]
1240 fn cancelling_dataset_weights_keep_their_squared_weight_uncertainty() {
1241 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1242 let cancelling = Dataset::from_events(
1243 schema,
1244 [
1245 OwnedEvent::weighted(vec![], vec![0.5], 1.0),
1246 OwnedEvent::weighted(vec![], vec![0.5], -1.0),
1247 ],
1248 )
1249 .unwrap();
1250 let histogram = cancelling
1251 .histogram(
1252 &event_scalar("x"),
1253 BinSpec::edges([0.0, 1.0]).unwrap(),
1254 true,
1255 None,
1256 &Execution::default(),
1257 )
1258 .unwrap();
1259
1260 assert_eq!(histogram.counts(), [0.0]);
1261 assert_eq!(histogram.sum_squared_weights(), [2.0]);
1262 assert_eq!(histogram.errors(), [2.0_f64.sqrt()]);
1263 }
1264
1265 #[test]
1266 fn empty_batches_are_valid_query_inputs() {
1267 let execution = Execution::default();
1268 let x = event_scalar("x");
1269
1270 let empty_batch_schema =
1271 Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1272 let empty_batch = Dataset::from_batch(
1273 EventBatch::from_events(empty_batch_schema, std::iter::empty::<OwnedEvent>()).unwrap(),
1274 );
1275 for empty in [empty_batch, dataset().empty_derived().unwrap()] {
1276 assert!(empty.evaluate_real(&x, &execution).unwrap().is_empty());
1277 assert!(
1278 empty
1279 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1280 .unwrap()
1281 .map_events(|event| event.scalar(0))
1282 .unwrap()
1283 .is_empty()
1284 );
1285 }
1286 }
1287
1288 #[test]
1289 fn event_column_nan_comparisons_are_false() {
1290 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1291 let dataset = Dataset::from_events(
1292 schema,
1293 [
1294 OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0),
1295 OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1296 OwnedEvent::weighted(vec![], vec![1.0], 1.0),
1297 ],
1298 )
1299 .unwrap();
1300 let x = event_scalar("x");
1301 let selected = dataset
1302 .select(&Predicate::ne(x, 0.0), &Execution::default())
1303 .unwrap();
1304
1305 assert_eq!(
1306 selected.map_events(|event| event.scalar(0)).unwrap(),
1307 vec![1.0]
1308 );
1309 }
1310
1311 #[test]
1312 fn query_propagates_source_batch_errors() {
1313 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1314 let dataset = Dataset::new(FailingSource { schema });
1315
1316 let error = dataset
1317 .evaluate_real(&event_scalar("x"), &Execution::default())
1318 .unwrap_err();
1319 assert!(
1320 matches!(error, RuntimeError::Data(message) if message.contains("query source failed"))
1321 );
1322 }
1323
1324 #[test]
1325 fn all_empty_bins_retain_valid_empty_derived_sources() {
1326 let source = dataset();
1327 let before = capability_tuple(source.capabilities());
1328 let bins = source
1329 .bin_by(
1330 &event_scalar("x"),
1331 BinSpec::edges([10.0, 20.0, 30.0]).unwrap(),
1332 &Execution::default(),
1333 )
1334 .unwrap();
1335
1336 assert_eq!(capability_tuple(source.capabilities()), before);
1337 assert_eq!(bins.len(), 2);
1338 for bin in bins {
1339 assert_eq!(bin.dataset().num_events().unwrap(), Some(0));
1340 assert!(
1341 bin.dataset()
1342 .evaluate_real(&event_scalar("x"), &Execution::default())
1343 .unwrap()
1344 .is_empty()
1345 );
1346 }
1347 }
1348
1349 #[test]
1350 fn traversing_all_bins_reads_the_source_once() {
1351 let reads = Arc::new(AtomicUsize::new(0));
1352 let source = CountingSource {
1353 inner: match dataset().batches().unwrap().next().unwrap() {
1354 Ok(batch) => MemorySource::new(batch),
1355 Err(error) => panic!("unexpected source error: {error}"),
1356 },
1357 reads: Arc::clone(&reads),
1358 };
1359 let dataset = Dataset::new(source).chunked(1).unwrap();
1360 let bins = dataset
1361 .bin_by(
1362 &event_scalar("x"),
1363 BinSpec::uniform(4, -1.0, 3.0).unwrap(),
1364 &Execution::default(),
1365 )
1366 .unwrap();
1367
1368 let values = bins
1369 .into_iter()
1370 .map(|bin| {
1371 bin.into_dataset()
1372 .map_events(|event| event.scalar(0))
1373 .unwrap()
1374 })
1375 .collect::<Vec<_>>();
1376 assert_eq!(values, [vec![-1.0], vec![0.0], vec![1.0], vec![2.0]]);
1377 assert_eq!(reads.load(Ordering::Relaxed), 1);
1378 }
1379
1380 #[test]
1381 fn between_predicates_have_explicit_endpoint_semantics() {
1382 let dataset = dataset();
1383 let execution = Execution::default();
1384 let x = event_scalar("x");
1385
1386 let closed = dataset
1387 .select(&Predicate::between(x.clone(), 0.0, 1.0), &execution)
1388 .unwrap();
1389 assert_eq!(
1390 closed.map_events(|event| event.scalar(0)).unwrap(),
1391 vec![0.0, 1.0]
1392 );
1393
1394 let open = dataset
1395 .select(
1396 &Predicate::between_with(x, -1.0, 1.0, IntervalClosure::Open),
1397 &execution,
1398 )
1399 .unwrap();
1400 assert_eq!(open.map_events(|event| event.scalar(0)).unwrap(), vec![0.0]);
1401 }
1402
1403 #[test]
1404 fn real_queries_reject_complex_and_free_parameter_expressions() {
1405 let dataset = dataset();
1406 let execution = Execution::default();
1407 assert!(
1408 dataset
1409 .evaluate_real(&complex(1.0, 1.0), &execution)
1410 .is_err()
1411 );
1412 let parameter = Expr::from(laddu_expr::parameters::Parameter::free("p"));
1413 assert!(dataset.evaluate_expr(¶meter, &execution).is_err());
1414 }
1415
1416 #[test]
1417 fn compiled_query_outputs_preserve_order_and_values() {
1418 let source = dataset();
1419 let batch = source.batches().unwrap().next().unwrap().unwrap();
1420 let x = event_scalar("x");
1421 let query = QueryExprSet::prepare(
1422 vec![x.clone() + 1.0, x.clone() * 2.0, x],
1423 &Execution::default(),
1424 false,
1425 )
1426 .unwrap();
1427 let values = query.evaluate_batch(&batch).unwrap();
1428 assert_eq!(
1429 values[0].iter().map(|v| v.re).collect::<Vec<_>>(),
1430 [0.0, 1.0, 2.0, 3.0]
1431 );
1432 assert_eq!(
1433 values[1].iter().map(|v| v.re).collect::<Vec<_>>(),
1434 [-2.0, 0.0, 2.0, 4.0]
1435 );
1436 assert_eq!(
1437 values[2].iter().map(|v| v.re).collect::<Vec<_>>(),
1438 [-1.0, 0.0, 1.0, 2.0]
1439 );
1440 }
1441
1442 #[test]
1443 fn repeated_predicate_leaves_are_evaluated_once() {
1444 let x = event_scalar("x");
1445 let selected = dataset()
1446 .select(
1447 &Predicate::ge(x.clone() + 1.0, 0.0).and(Predicate::lt(x + 1.0, 2.0)),
1448 &Execution::default(),
1449 )
1450 .unwrap();
1451 assert_eq!(
1452 selected.map_events(|event| event.scalar(0)).unwrap(),
1453 [-1.0, 0.0]
1454 );
1455 }
1456
1457 #[test]
1458 fn f32_queries_match_f64_query_results() {
1459 let x = event_scalar("x");
1460 let f64_values = dataset().evaluate_real(&x, &Execution::default()).unwrap();
1461 let f32_execution = Execution::local(ExecutionOptions {
1462 device: Device::Cpu(CpuOptions::default()),
1463 precision: Precision::F32,
1464 ..ExecutionOptions::default()
1465 })
1466 .unwrap();
1467 let f32_values = dataset().evaluate_real(&x, &f32_execution).unwrap();
1468 assert_eq!(f32_values, f64_values);
1469
1470 let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1471 let f64_histogram = dataset()
1472 .histogram(&x, bins.clone(), true, None, &Execution::default())
1473 .unwrap();
1474 let f32_histogram = dataset()
1475 .histogram(&x, bins, true, None, &f32_execution)
1476 .unwrap();
1477 assert_eq!(f32_histogram, f64_histogram);
1478 }
1479
1480 #[test]
1481 fn bin_edges_validate_and_nan_predicates_are_false() {
1482 assert!(BinSpec::edges([0.0, 0.0]).is_err());
1483 assert!(!compare(f64::NAN, Comparison::Ne, 0.0));
1484 }
1485
1486 #[test]
1487 fn selection_is_lazy_and_one_pass_binning_preserves_streaming_policy() {
1488 let source = dataset();
1489 let batch = source.batches().unwrap().next().unwrap().unwrap();
1490 let reads = Arc::new(AtomicUsize::new(0));
1491 let dataset = Dataset::new(CountingSource {
1492 inner: MemorySource::new(batch),
1493 reads: Arc::clone(&reads),
1494 })
1495 .streaming();
1496 let execution = Execution::default();
1497 let x = event_scalar("x");
1498
1499 let selected = dataset
1500 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1501 .unwrap();
1502 let bins = dataset
1503 .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1504 .unwrap();
1505 assert_eq!(reads.load(Ordering::Relaxed), 1);
1506 assert_eq!(
1507 selected.cache_storage(),
1508 laddu_data::data::CacheStorage::Streaming
1509 );
1510
1511 assert_eq!(
1512 selected.map_events(|event| event.scalar(0)).unwrap(),
1513 vec![0.0, 1.0, 2.0]
1514 );
1515 assert_eq!(reads.load(Ordering::Relaxed), 2);
1516 assert_eq!(
1517 bins[0]
1518 .dataset()
1519 .map_events(|event| event.scalar(0))
1520 .unwrap(),
1521 vec![0.0]
1522 );
1523 assert_eq!(reads.load(Ordering::Relaxed), 2);
1524 }
1525
1526 #[test]
1527 fn unknown_cardinality_fastest_discovers_and_retains_small_selection() {
1528 let source = dataset();
1529 let batch = source.batches().unwrap().next().unwrap().unwrap();
1530 let reads = Arc::new(AtomicUsize::new(0));
1531 let dataset = Dataset::new(CountingSource {
1532 inner: MemorySource::new(batch),
1533 reads: Arc::clone(&reads),
1534 });
1535 let execution = Execution::default();
1536 let x = event_scalar("x");
1537 let selected = dataset
1538 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1539 .unwrap();
1540 let compiled = CompiledModel::from_expr(&x).unwrap();
1541 let params = compiled.params().default_values();
1542 let model = PreparedModel::prepare(&compiled, &execution).unwrap();
1543 let prepared = model.prepare_dataset(&execution, &selected).unwrap();
1544
1545 #[cfg(not(feature = "wgpu"))]
1546 let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared;
1547 #[cfg(feature = "wgpu")]
1548 let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared else {
1549 panic!("default execution prepares CPU datasets");
1550 };
1551 assert_eq!(
1552 prepared_cpu.stats().storage(),
1553 laddu_data::data::CacheStorage::Resident
1554 );
1555 assert_eq!(prepared_cpu.stats().local_events(), 3);
1556 assert_eq!(reads.load(Ordering::Relaxed), 2);
1557
1558 for _ in 0..2 {
1559 assert_eq!(
1560 model
1561 .reduce(
1562 &execution,
1563 ¶ms,
1564 &prepared,
1565 laddu_compile::ReductionPlan::weighted_real(),
1566 )
1567 .unwrap(),
1568 5.5
1569 );
1570 }
1571 assert_eq!(reads.load(Ordering::Relaxed), 2);
1572 }
1573}