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 PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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 = PreparedQuery::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
542pub struct PreparedQuery {
545 shared: QueryExprStorage,
546 outputs: usize,
547}
548
549enum QueryExprStorage {
550 Shared(QueryExpr),
551 Separate(Vec<QueryExpr>),
552}
553
554impl PreparedQuery {
555 pub fn prepare(
561 expressions: Vec<Expr>,
562 execution: &Execution,
563 require_real: bool,
564 ) -> RuntimeResult<Self> {
565 let expression_count = expressions.len();
566 let compiled = CompiledQuery::from_exprs(expressions.clone())
567 .map_err(|error| query_error(error.to_string()))?;
568 let model = compiled.model();
569 let outputs = compiled.outputs();
570 if outputs.len() != expression_count {
571 return Err(query_error(
572 "compiled query output count changed during lowering",
573 ));
574 }
575 for element in outputs {
576 let value_kind = model
577 .node_facts(*element)
578 .map(|facts| facts.value_kind)
579 .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
580 if value_kind != ValueKind::Real
581 && (require_real || !matches!(value_kind, ValueKind::Complex))
582 {
583 return Err(query_error(if require_real {
584 "this dataset operation requires a real-valued expression"
585 } else {
586 "dataset expressions must be scalar"
587 }));
588 }
589 }
590 if model.params().n_free() != 0 {
591 return Err(query_error(
592 "dataset expressions cannot contain free parameters",
593 ));
594 }
595 let params = model.params().default_values();
596 let plan = match PreparedModel::prepare(model, execution) {
597 Ok(plan) => QueryExprStorage::Shared(QueryExpr {
598 model: plan,
599 params,
600 outputs: outputs.to_vec(),
601 }),
602 Err(shared_error) => {
603 if !may_fallback_to_scalar(execution, &shared_error) {
604 return Err(shared_error);
605 }
606 let separate = expressions
607 .iter()
608 .map(|expr| QueryExpr::prepare(expr, execution, require_real))
609 .collect::<RuntimeResult<Vec<_>>>();
610 QueryExprStorage::Separate(separate?)
611 }
612 };
613 Ok(Self {
614 shared: plan,
615 outputs: expression_count,
616 })
617 }
618
619 #[doc(hidden)]
621 pub fn batch_memory_estimate(&self, events: usize) -> usize {
622 let model = match &self.shared {
623 QueryExprStorage::Shared(query) => query.model.batch_memory_estimate(events),
624 QueryExprStorage::Separate(queries) => queries
625 .iter()
626 .map(|query| query.model.batch_memory_estimate(events))
627 .max()
628 .unwrap_or(0),
629 };
630 model.saturating_add(events.saturating_mul(self.outputs).saturating_mul(32))
631 }
632
633 pub fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Vec<Complex64>>> {
638 let values = match &self.shared {
639 QueryExprStorage::Shared(query) => {
640 query
641 .model
642 .evaluate_batch_outputs(&query.params, batch, &query.outputs)?
643 }
644 QueryExprStorage::Separate(queries) => queries
645 .iter()
646 .map(|query| query.evaluate_batch(batch))
647 .collect::<RuntimeResult<Vec<_>>>()?,
648 };
649 if values.len() != self.outputs {
650 return Err(query_error(
651 "compiled query returned an unexpected output count",
652 ));
653 }
654 Ok(values)
655 }
656}
657
658impl QueryExpr {
659 fn prepare(expr: &Expr, execution: &Execution, require_real: bool) -> RuntimeResult<Self> {
660 if expr.shape().map_err(|e| query_error(e.to_string()))? != laddu_expr::ExprShape::Scalar {
661 return Err(query_error("dataset expressions must be scalar"));
662 }
663 let compiled = laddu_compile::CompiledModel::from_expr(expr)
664 .map_err(|error| query_error(error.to_string()))?;
665 let value_kind = compiled
666 .node_facts(compiled.graph().root())
667 .map(|facts| facts.value_kind)
668 .ok_or_else(|| query_error("compiled expression facts are incomplete"))?;
669 if require_real && value_kind == ValueKind::Complex {
670 return Err(query_error(
671 "this dataset operation requires a real-valued expression",
672 ));
673 }
674 if compiled.params().n_free() != 0 {
675 return Err(query_error(
676 "dataset expressions cannot contain free parameters",
677 ));
678 }
679 let params = compiled.params().default_values();
680 let model = PreparedModel::prepare(&compiled, execution)?;
681 Ok(Self {
682 model,
683 params,
684 outputs: Vec::new(),
685 })
686 }
687
688 fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<Complex64>> {
689 self.model.evaluate_batch(&self.params, batch)
690 }
691}
692
693struct CompiledPredicate {
694 expressions: PreparedQuery,
695 program: PredicateProgram,
696}
697
698enum PredicateProgram {
699 Compare {
700 lhs: usize,
701 op: Comparison,
702 rhs: usize,
703 },
704 And(Box<Self>, Box<Self>),
705 Or(Box<Self>, Box<Self>),
706 Not(Box<Self>),
707 Between {
708 value: usize,
709 lower: usize,
710 upper: usize,
711 closure: IntervalClosure,
712 },
713}
714
715impl CompiledPredicate {
716 fn prepare(predicate: &Predicate, execution: &Execution) -> RuntimeResult<Self> {
717 let mut expressions = Vec::new();
718 let program = Self::compile_program(predicate, &mut expressions);
719 Ok(Self {
720 expressions: PreparedQuery::prepare(expressions, execution, true)?,
721 program,
722 })
723 }
724
725 fn compile_program(predicate: &Predicate, expressions: &mut Vec<Expr>) -> PredicateProgram {
726 let leaf = |expr: &Expr, expressions: &mut Vec<Expr>| {
727 let index = expressions.len();
728 expressions.push(expr.clone());
729 index
730 };
731 match predicate {
732 Predicate::Compare { lhs, op, rhs } => PredicateProgram::Compare {
733 lhs: leaf(lhs, expressions),
734 op: *op,
735 rhs: leaf(rhs, expressions),
736 },
737 Predicate::And(lhs, rhs) => PredicateProgram::And(
738 Box::new(Self::compile_program(lhs, expressions)),
739 Box::new(Self::compile_program(rhs, expressions)),
740 ),
741 Predicate::Or(lhs, rhs) => PredicateProgram::Or(
742 Box::new(Self::compile_program(lhs, expressions)),
743 Box::new(Self::compile_program(rhs, expressions)),
744 ),
745 Predicate::Not(inner) => {
746 PredicateProgram::Not(Box::new(Self::compile_program(inner, expressions)))
747 }
748 Predicate::Between {
749 value,
750 lower,
751 upper,
752 closure,
753 } => PredicateProgram::Between {
754 value: leaf(value, expressions),
755 lower: leaf(lower, expressions),
756 upper: leaf(upper, expressions),
757 closure: *closure,
758 },
759 }
760 }
761
762 fn evaluate_batch(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
763 let values = self.expressions.evaluate_batch(batch)?;
764 Ok((0..batch.len())
765 .filter(|row| Self::evaluate_row(&self.program, &values, *row))
766 .collect())
767 }
768
769 fn evaluate_row(program: &PredicateProgram, values: &[Vec<Complex64>], row: usize) -> bool {
770 match program {
771 PredicateProgram::Compare { lhs, op, rhs } => {
772 compare(values[*lhs][row].re, *op, values[*rhs][row].re)
773 }
774 PredicateProgram::And(lhs, rhs) => {
775 Self::evaluate_row(lhs, values, row) && Self::evaluate_row(rhs, values, row)
776 }
777 PredicateProgram::Or(lhs, rhs) => {
778 Self::evaluate_row(lhs, values, row) || Self::evaluate_row(rhs, values, row)
779 }
780 PredicateProgram::Not(inner) => !Self::evaluate_row(inner, values, row),
781 PredicateProgram::Between {
782 value,
783 lower,
784 upper,
785 closure,
786 } => {
787 let lower_op = match closure {
788 IntervalClosure::Open | IntervalClosure::RightClosed => Comparison::Gt,
789 IntervalClosure::LeftClosed | IntervalClosure::Closed => Comparison::Ge,
790 };
791 let upper_op = match closure {
792 IntervalClosure::Open | IntervalClosure::LeftClosed => Comparison::Lt,
793 IntervalClosure::RightClosed | IntervalClosure::Closed => Comparison::Le,
794 };
795 compare(values[*value][row].re, lower_op, values[*lower][row].re)
796 && compare(values[*value][row].re, upper_op, values[*upper][row].re)
797 }
798 }
799 }
800}
801
802fn compare(lhs: f64, op: Comparison, rhs: f64) -> bool {
803 if lhs.is_nan() || rhs.is_nan() {
804 return false;
805 }
806 match op {
807 Comparison::Lt => lhs < rhs,
808 Comparison::Le => lhs <= rhs,
809 Comparison::Gt => lhs > rhs,
810 Comparison::Ge => lhs >= rhs,
811 Comparison::Eq => lhs == rhs,
812 Comparison::Ne => lhs != rhs,
813 }
814}
815
816#[derive(Clone)]
817struct QuerySource {
818 source: Dataset,
819 filter: QueryFilter,
820}
821
822#[derive(Clone)]
823enum QueryFilter {
824 Predicate(Arc<CompiledPredicate>),
825}
826
827impl EventSource for QuerySource {
828 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
829 self.source.schema()
830 }
831
832 fn capabilities(&self) -> SourceCapabilities {
833 let source = self.source.capabilities();
834 SourceCapabilities {
835 exact_len: false,
836 exact_weighted_total: false,
837 random_access: false,
838 deterministic_partitioning: source.deterministic_partitioning,
839 predicate_pushdown: false,
840 projection_pushdown: false,
841 streaming: true,
842 }
843 }
844
845 fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
846 let batches = self.source.stream_with_plan(plan)?;
847 let filter = self.filter.clone();
848 Ok(Box::new(batches.filter_map(move |batch| {
849 let batch = match batch {
850 Ok(batch) => batch,
851 Err(error) => return Some(Err(error)),
852 };
853 let rows = match filter.rows(&batch) {
854 Ok(rows) => rows,
855 Err(error) => return Some(Err(LadduDataError::Source(error.to_string()))),
856 };
857 (!rows.is_empty()).then(|| Ok(batch.select(&rows)))
858 })))
859 }
860}
861
862impl QueryFilter {
863 fn rows(&self, batch: &EventBatch) -> RuntimeResult<Vec<usize>> {
864 match self {
865 Self::Predicate(predicate) => predicate.evaluate_batch(batch),
866 }
867 }
868}
869
870fn query_error(message: impl Into<String>) -> RuntimeError {
871 RuntimeError::InvalidShape {
872 index: 0,
873 message: message.into(),
874 }
875}
876
877fn may_fallback_to_scalar(execution: &Execution, error: &RuntimeError) -> bool {
878 let cpu_f32 = matches!(
879 error,
880 RuntimeError::Execution(crate::ExecutionError::UnsupportedCpuF32Model)
881 );
882 #[cfg(feature = "wgpu")]
883 {
884 cpu_f32 || (execution.wgpu_context().is_some() && matches!(error, RuntimeError::Wgpu(_)))
885 }
886 #[cfg(not(feature = "wgpu"))]
887 {
888 let _ = execution;
889 cpu_f32
890 }
891}
892
893fn data_error(error: impl ToString) -> RuntimeError {
894 RuntimeError::Data(error.to_string())
895}
896
897#[cfg(test)]
898mod tests {
899 use super::*;
900 use crate::{CpuOptions, Device, ExecutionOptions, Precision};
901 #[cfg(feature = "jit")]
902 use crate::{JitPolicy, ThreadPolicy};
903 use laddu_compile::CompiledModel;
904 use laddu_data::{
905 data::{EventBatch, OwnedEvent},
906 io::{EventSource, ReadPlan, SourceCapabilities, memory::MemorySource},
907 schema::Schema,
908 };
909 use laddu_expr::{complex, event_scalar};
910 use std::sync::atomic::{AtomicUsize, Ordering};
911
912 #[test]
913 fn bin_spec_roundtrip_preserves_validation() {
914 let bins = BinSpec::edges([-1.0, 0.0, 2.0]).unwrap();
915 let json = serde_json::to_string(&bins).unwrap();
916 assert_eq!(serde_json::from_str::<BinSpec>(&json).unwrap(), bins);
917 assert!(serde_json::from_str::<BinSpec>("[0.0,0.0]").is_err());
918 }
919
920 #[derive(Clone)]
921 struct CountingSource {
922 inner: MemorySource,
923 reads: Arc<AtomicUsize>,
924 }
925
926 impl EventSource for CountingSource {
927 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
928 EventSource::schema(&self.inner)
929 }
930
931 fn capabilities(&self) -> SourceCapabilities {
932 self.inner.capabilities()
933 }
934
935 fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
936 self.reads.fetch_add(1, Ordering::Relaxed);
937 self.inner.batches(plan)
938 }
939 }
940
941 #[derive(Clone)]
942 struct FailingSource {
943 schema: Arc<Schema>,
944 }
945
946 impl EventSource for FailingSource {
947 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
948 Ok(Arc::clone(&self.schema))
949 }
950
951 fn capabilities(&self) -> SourceCapabilities {
952 SourceCapabilities {
953 exact_len: false,
954 exact_weighted_total: false,
955 random_access: false,
956 deterministic_partitioning: true,
957 predicate_pushdown: false,
958 projection_pushdown: false,
959 streaming: true,
960 }
961 }
962
963 fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
964 Ok(Box::new(std::iter::once(Err(LadduDataError::Source(
965 "query source failed".into(),
966 )))))
967 }
968 }
969
970 fn capability_tuple(
971 capabilities: SourceCapabilities,
972 ) -> (bool, bool, bool, bool, bool, bool, bool) {
973 (
974 capabilities.exact_len,
975 capabilities.exact_weighted_total,
976 capabilities.random_access,
977 capabilities.deterministic_partitioning,
978 capabilities.predicate_pushdown,
979 capabilities.projection_pushdown,
980 capabilities.streaming,
981 )
982 }
983
984 fn dataset() -> Dataset {
985 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
986 Dataset::from_events(
987 schema,
988 [
989 OwnedEvent::weighted(vec![], vec![-1.0], 0.5),
990 OwnedEvent::weighted(vec![], vec![0.0], 1.0),
991 OwnedEvent::weighted(vec![], vec![1.0], 1.5),
992 OwnedEvent::weighted(vec![], vec![2.0], 2.0),
993 ],
994 )
995 .unwrap()
996 }
997
998 #[test]
999 fn evaluates_selects_and_bins_dataset_expressions() {
1000 let dataset = dataset().chunked(1).unwrap();
1001 let execution = Execution::default();
1002 let x = event_scalar("x");
1003 assert_eq!(
1004 dataset.evaluate_real(&x, &execution).unwrap(),
1005 vec![-1.0, 0.0, 1.0, 2.0]
1006 );
1007
1008 let selected = dataset
1009 .select(
1010 &Predicate::ge(x.clone(), 0.0).and(Predicate::lt(x.clone(), 2.0)),
1011 &execution,
1012 )
1013 .unwrap();
1014 assert_eq!(
1015 selected.map_events(|event| event.scalar(0)).unwrap(),
1016 vec![0.0, 1.0]
1017 );
1018 assert_eq!(selected.sum_weights().unwrap(), 2.5);
1019
1020 let bins = dataset
1021 .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1022 .unwrap();
1023 assert_eq!(bins.len(), 2);
1024 assert_eq!(
1025 bins[0]
1026 .dataset()
1027 .map_events(|event| event.scalar(0))
1028 .unwrap(),
1029 vec![0.0]
1030 );
1031 assert_eq!(
1032 bins[1]
1033 .dataset()
1034 .map_events(|event| event.scalar(0))
1035 .unwrap(),
1036 vec![1.0, 2.0]
1037 );
1038 }
1039
1040 #[test]
1041 fn dataset_histogram_uses_event_weights_and_excludes_the_final_upper_edge() {
1042 let histogram = dataset()
1043 .histogram(
1044 &event_scalar("x"),
1045 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1046 true,
1047 None,
1048 &Execution::default(),
1049 )
1050 .unwrap();
1051
1052 assert_eq!(histogram.counts(), [1.0, 1.5]);
1053 assert_eq!(histogram.sum_squared_weights(), [1.0, 2.25]);
1054 assert_eq!(histogram.underflow(), 0.5);
1055 assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.25));
1056 assert_eq!(histogram.overflow(), 2.0);
1057 assert_eq!(histogram.overflow_sum_squared_weights(), Some(4.0));
1058 }
1059
1060 #[test]
1061 fn dataset_joint_histogram_uses_row_major_bins_and_aggregates_invalid_events() {
1062 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x", "y"], true).unwrap());
1063 let dataset = Dataset::from_events(
1064 schema,
1065 [
1066 OwnedEvent::weighted(vec![], vec![0.0, 10.0], 1.0),
1067 OwnedEvent::weighted(vec![], vec![1.0, 10.0], -2.0),
1068 OwnedEvent::weighted(vec![], vec![0.0, 20.0], 3.0),
1069 OwnedEvent::weighted(vec![], vec![f64::NAN, 10.0], 4.0),
1070 ],
1071 )
1072 .unwrap();
1073
1074 let histogram = dataset
1075 .joint_histogram(
1076 &[event_scalar("x"), event_scalar("y")],
1077 vec![
1078 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1079 BinSpec::edges([0.0, 15.0, 25.0]).unwrap(),
1080 ],
1081 true,
1082 None,
1083 &Execution::default(),
1084 )
1085 .unwrap();
1086
1087 assert_eq!(histogram.shape(), [2, 2]);
1088 assert_eq!(histogram.values(), [1.0, 3.0, -2.0, 0.0]);
1089 assert_eq!(histogram.sum_squared_weights(), [1.0, 9.0, 4.0, 0.0]);
1090 assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
1091 assert_eq!(histogram.diagnostics().out_of_range_count(), 0);
1092 }
1093
1094 #[test]
1095 fn dataset_histogram_multiplies_the_selected_base_and_custom_weights() {
1096 let x = event_scalar("x");
1097 let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1098 let execution = Execution::default();
1099
1100 let weighted = dataset()
1101 .histogram(&x, bins.clone(), true, Some(&(x.clone() + 2.0)), &execution)
1102 .unwrap();
1103 assert_eq!(weighted.counts(), [0.5, 6.5]);
1104 assert_eq!(weighted.sum_squared_weights(), [0.25, 24.25]);
1105 assert_eq!(weighted.overflow(), 8.0);
1106
1107 let custom_only = dataset()
1108 .histogram(&x, bins, false, Some(&(x.clone() + 2.0)), &execution)
1109 .unwrap();
1110 assert_eq!(custom_only.counts(), [1.0, 5.0]);
1111 assert_eq!(custom_only.sum_squared_weights(), [1.0, 13.0]);
1112 assert_eq!(custom_only.overflow(), 4.0);
1113 }
1114
1115 #[test]
1116 fn dataset_histogram_preserves_view_source_and_memory_semantics() {
1117 let source = dataset();
1118 let batch = source.batches().unwrap().next().unwrap().unwrap();
1119 let reads = Arc::new(AtomicUsize::new(0));
1120 let counted = Dataset::new(CountingSource {
1121 inner: MemorySource::new(batch),
1122 reads: Arc::clone(&reads),
1123 })
1124 .streaming()
1125 .chunked(1)
1126 .unwrap();
1127 let selected = counted
1128 .select(
1129 &Predicate::ge(event_scalar("x"), 0.0),
1130 &Execution::default(),
1131 )
1132 .unwrap();
1133
1134 let histogram = selected
1135 .histogram(
1136 &event_scalar("x"),
1137 BinSpec::edges([0.0, 1.0, 2.0]).unwrap(),
1138 true,
1139 None,
1140 &Execution::default(),
1141 )
1142 .unwrap();
1143
1144 assert_eq!(histogram.counts(), [1.0, 1.5]);
1145 assert_eq!(histogram.overflow(), 2.0);
1146 assert_eq!(reads.load(Ordering::Relaxed), 1);
1147 }
1148
1149 #[test]
1150 fn dataset_histogram_matches_across_memory_policies_and_chunking() {
1151 let source = dataset();
1152 let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1153 let execution = Execution::default();
1154 let expected = source
1155 .clone()
1156 .resident()
1157 .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1158 .unwrap();
1159
1160 for candidate in [
1161 source.clone().streaming(),
1162 source.clone().resident().chunked(1).unwrap(),
1163 source.streaming().chunked(2).unwrap(),
1164 ] {
1165 let actual = candidate
1166 .histogram(&event_scalar("x"), bins.clone(), true, None, &execution)
1167 .unwrap();
1168 assert_eq!(actual, expected);
1169 }
1170 }
1171
1172 #[cfg(feature = "jit")]
1173 #[test]
1174 fn dataset_histogram_matches_cpu_interpreter_and_jit_backends() {
1175 let source = dataset();
1176 let bins = BinSpec::edges([0.0, 1.0, 2.0]).unwrap();
1177 let execution = |jit| {
1178 Execution::local(ExecutionOptions {
1179 device: Device::Cpu(CpuOptions {
1180 threads: ThreadPolicy::Serial,
1181 jit,
1182 }),
1183 precision: Precision::F64,
1184 ..ExecutionOptions::default()
1185 })
1186 .unwrap()
1187 };
1188 let interpreted = source
1189 .histogram(
1190 &event_scalar("x"),
1191 bins.clone(),
1192 true,
1193 None,
1194 &execution(JitPolicy::Disabled),
1195 )
1196 .unwrap();
1197 let compiled = source
1198 .histogram(
1199 &event_scalar("x"),
1200 bins,
1201 true,
1202 None,
1203 &execution(JitPolicy::Enabled),
1204 )
1205 .unwrap();
1206
1207 assert_eq!(compiled, interpreted);
1208 }
1209
1210 #[test]
1211 fn dataset_histogram_reports_invalid_values_and_source_failures() {
1212 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1213 let nonfinite = Dataset::from_events(
1214 Arc::clone(&schema),
1215 [OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0)],
1216 )
1217 .unwrap();
1218 let error = nonfinite
1219 .histogram(
1220 &event_scalar("x"),
1221 BinSpec::edges([0.0, 1.0]).unwrap(),
1222 true,
1223 None,
1224 &Execution::default(),
1225 )
1226 .unwrap_err();
1227 assert!(
1228 error.to_string().contains("expected finite, got NaN"),
1229 "unexpected error: {error}",
1230 );
1231
1232 let source_error = Dataset::new(FailingSource { schema })
1233 .histogram(
1234 &event_scalar("x"),
1235 BinSpec::edges([0.0, 1.0]).unwrap(),
1236 true,
1237 None,
1238 &Execution::default(),
1239 )
1240 .unwrap_err();
1241 assert!(source_error.to_string().contains("query source failed"));
1242 }
1243
1244 #[test]
1245 fn empty_dataset_histogram_is_a_valid_empirical_histogram() {
1246 let histogram = dataset()
1247 .empty_derived()
1248 .unwrap()
1249 .histogram(
1250 &event_scalar("x"),
1251 BinSpec::edges([0.0, 1.0, 3.0]).unwrap(),
1252 true,
1253 None,
1254 &Execution::default(),
1255 )
1256 .unwrap();
1257
1258 assert_eq!(histogram.counts(), [0.0, 0.0]);
1259 assert_eq!(histogram.sum_squared_weights(), [0.0, 0.0]);
1260 assert_eq!(histogram.underflow_sum_squared_weights(), Some(0.0));
1261 assert_eq!(histogram.overflow_sum_squared_weights(), Some(0.0));
1262 }
1263
1264 #[test]
1265 fn cancelling_dataset_weights_keep_their_squared_weight_uncertainty() {
1266 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1267 let cancelling = Dataset::from_events(
1268 schema,
1269 [
1270 OwnedEvent::weighted(vec![], vec![0.5], 1.0),
1271 OwnedEvent::weighted(vec![], vec![0.5], -1.0),
1272 ],
1273 )
1274 .unwrap();
1275 let histogram = cancelling
1276 .histogram(
1277 &event_scalar("x"),
1278 BinSpec::edges([0.0, 1.0]).unwrap(),
1279 true,
1280 None,
1281 &Execution::default(),
1282 )
1283 .unwrap();
1284
1285 assert_eq!(histogram.counts(), [0.0]);
1286 assert_eq!(histogram.sum_squared_weights(), [2.0]);
1287 assert_eq!(histogram.errors(), [2.0_f64.sqrt()]);
1288 }
1289
1290 #[test]
1291 fn empty_batches_are_valid_query_inputs() {
1292 let execution = Execution::default();
1293 let x = event_scalar("x");
1294
1295 let empty_batch_schema =
1296 Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1297 let empty_batch = Dataset::from_batch(
1298 EventBatch::from_events(empty_batch_schema, std::iter::empty::<OwnedEvent>()).unwrap(),
1299 );
1300 for empty in [empty_batch, dataset().empty_derived().unwrap()] {
1301 assert!(empty.evaluate_real(&x, &execution).unwrap().is_empty());
1302 assert!(
1303 empty
1304 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1305 .unwrap()
1306 .map_events(|event| event.scalar(0))
1307 .unwrap()
1308 .is_empty()
1309 );
1310 }
1311 }
1312
1313 #[test]
1314 fn event_column_nan_comparisons_are_false() {
1315 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1316 let dataset = Dataset::from_events(
1317 schema,
1318 [
1319 OwnedEvent::weighted(vec![], vec![f64::NAN], 1.0),
1320 OwnedEvent::weighted(vec![], vec![0.0], 1.0),
1321 OwnedEvent::weighted(vec![], vec![1.0], 1.0),
1322 ],
1323 )
1324 .unwrap();
1325 let x = event_scalar("x");
1326 let selected = dataset
1327 .select(&Predicate::ne(x, 0.0), &Execution::default())
1328 .unwrap();
1329
1330 assert_eq!(
1331 selected.map_events(|event| event.scalar(0)).unwrap(),
1332 vec![1.0]
1333 );
1334 }
1335
1336 #[test]
1337 fn query_propagates_source_batch_errors() {
1338 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1339 let dataset = Dataset::new(FailingSource { schema });
1340
1341 let error = dataset
1342 .evaluate_real(&event_scalar("x"), &Execution::default())
1343 .unwrap_err();
1344 assert!(
1345 matches!(error, RuntimeError::Data(message) if message.contains("query source failed"))
1346 );
1347 }
1348
1349 #[test]
1350 fn all_empty_bins_retain_valid_empty_derived_sources() {
1351 let source = dataset();
1352 let before = capability_tuple(source.capabilities());
1353 let bins = source
1354 .bin_by(
1355 &event_scalar("x"),
1356 BinSpec::edges([10.0, 20.0, 30.0]).unwrap(),
1357 &Execution::default(),
1358 )
1359 .unwrap();
1360
1361 assert_eq!(capability_tuple(source.capabilities()), before);
1362 assert_eq!(bins.len(), 2);
1363 for bin in bins {
1364 assert_eq!(bin.dataset().num_events().unwrap(), Some(0));
1365 assert!(
1366 bin.dataset()
1367 .evaluate_real(&event_scalar("x"), &Execution::default())
1368 .unwrap()
1369 .is_empty()
1370 );
1371 }
1372 }
1373
1374 #[test]
1375 fn traversing_all_bins_reads_the_source_once() {
1376 let reads = Arc::new(AtomicUsize::new(0));
1377 let source = CountingSource {
1378 inner: match dataset().batches().unwrap().next().unwrap() {
1379 Ok(batch) => MemorySource::new(batch),
1380 Err(error) => panic!("unexpected source error: {error}"),
1381 },
1382 reads: Arc::clone(&reads),
1383 };
1384 let dataset = Dataset::new(source).chunked(1).unwrap();
1385 let bins = dataset
1386 .bin_by(
1387 &event_scalar("x"),
1388 BinSpec::uniform(4, -1.0, 3.0).unwrap(),
1389 &Execution::default(),
1390 )
1391 .unwrap();
1392
1393 let values = bins
1394 .into_iter()
1395 .map(|bin| {
1396 bin.into_dataset()
1397 .map_events(|event| event.scalar(0))
1398 .unwrap()
1399 })
1400 .collect::<Vec<_>>();
1401 assert_eq!(values, [vec![-1.0], vec![0.0], vec![1.0], vec![2.0]]);
1402 assert_eq!(reads.load(Ordering::Relaxed), 1);
1403 }
1404
1405 #[test]
1406 fn between_predicates_have_explicit_endpoint_semantics() {
1407 let dataset = dataset();
1408 let execution = Execution::default();
1409 let x = event_scalar("x");
1410
1411 let closed = dataset
1412 .select(&Predicate::between(x.clone(), 0.0, 1.0), &execution)
1413 .unwrap();
1414 assert_eq!(
1415 closed.map_events(|event| event.scalar(0)).unwrap(),
1416 vec![0.0, 1.0]
1417 );
1418
1419 let open = dataset
1420 .select(
1421 &Predicate::between_with(x, -1.0, 1.0, IntervalClosure::Open),
1422 &execution,
1423 )
1424 .unwrap();
1425 assert_eq!(open.map_events(|event| event.scalar(0)).unwrap(), vec![0.0]);
1426 }
1427
1428 #[test]
1429 fn real_queries_reject_complex_and_free_parameter_expressions() {
1430 let dataset = dataset();
1431 let execution = Execution::default();
1432 assert!(
1433 dataset
1434 .evaluate_real(&complex(1.0, 1.0), &execution)
1435 .is_err()
1436 );
1437 let parameter = Expr::from(laddu_expr::parameters::Parameter::free("p"));
1438 assert!(dataset.evaluate_expr(¶meter, &execution).is_err());
1439 }
1440
1441 #[test]
1442 fn compiled_query_outputs_preserve_order_and_values() {
1443 let source = dataset();
1444 let batch = source.batches().unwrap().next().unwrap().unwrap();
1445 let x = event_scalar("x");
1446 let query = PreparedQuery::prepare(
1447 vec![x.clone() + 1.0, x.clone() * 2.0, x],
1448 &Execution::default(),
1449 false,
1450 )
1451 .unwrap();
1452 let values = query.evaluate_batch(&batch).unwrap();
1453 assert_eq!(
1454 values[0].iter().map(|v| v.re).collect::<Vec<_>>(),
1455 [0.0, 1.0, 2.0, 3.0]
1456 );
1457 assert_eq!(
1458 values[1].iter().map(|v| v.re).collect::<Vec<_>>(),
1459 [-2.0, 0.0, 2.0, 4.0]
1460 );
1461 assert_eq!(
1462 values[2].iter().map(|v| v.re).collect::<Vec<_>>(),
1463 [-1.0, 0.0, 1.0, 2.0]
1464 );
1465 }
1466
1467 #[test]
1468 fn repeated_predicate_leaves_are_evaluated_once() {
1469 let x = event_scalar("x");
1470 let selected = dataset()
1471 .select(
1472 &Predicate::ge(x.clone() + 1.0, 0.0).and(Predicate::lt(x + 1.0, 2.0)),
1473 &Execution::default(),
1474 )
1475 .unwrap();
1476 assert_eq!(
1477 selected.map_events(|event| event.scalar(0)).unwrap(),
1478 [-1.0, 0.0]
1479 );
1480 }
1481
1482 #[test]
1483 fn f32_queries_match_f64_query_results() {
1484 let x = event_scalar("x");
1485 let f64_values = dataset().evaluate_real(&x, &Execution::default()).unwrap();
1486 let f32_execution = Execution::local(ExecutionOptions {
1487 device: Device::Cpu(CpuOptions::default()),
1488 precision: Precision::F32,
1489 ..ExecutionOptions::default()
1490 })
1491 .unwrap();
1492 let f32_values = dataset().evaluate_real(&x, &f32_execution).unwrap();
1493 assert_eq!(f32_values, f64_values);
1494
1495 let bins = BinSpec::edges([-2.0, 0.0, 2.0]).unwrap();
1496 let f64_histogram = dataset()
1497 .histogram(&x, bins.clone(), true, None, &Execution::default())
1498 .unwrap();
1499 let f32_histogram = dataset()
1500 .histogram(&x, bins, true, None, &f32_execution)
1501 .unwrap();
1502 assert_eq!(f32_histogram, f64_histogram);
1503 }
1504
1505 #[test]
1506 fn bin_edges_validate_and_nan_predicates_are_false() {
1507 assert!(BinSpec::edges([0.0, 0.0]).is_err());
1508 assert!(!compare(f64::NAN, Comparison::Ne, 0.0));
1509 }
1510
1511 #[test]
1512 fn selection_is_lazy_and_one_pass_binning_preserves_streaming_policy() {
1513 let source = dataset();
1514 let batch = source.batches().unwrap().next().unwrap().unwrap();
1515 let reads = Arc::new(AtomicUsize::new(0));
1516 let dataset = Dataset::new(CountingSource {
1517 inner: MemorySource::new(batch),
1518 reads: Arc::clone(&reads),
1519 })
1520 .streaming();
1521 let execution = Execution::default();
1522 let x = event_scalar("x");
1523
1524 let selected = dataset
1525 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1526 .unwrap();
1527 let bins = dataset
1528 .bin_by(&x, BinSpec::uniform(2, 0.0, 2.0).unwrap(), &execution)
1529 .unwrap();
1530 assert_eq!(reads.load(Ordering::Relaxed), 1);
1531 assert_eq!(
1532 selected.cache_storage(),
1533 laddu_data::data::CacheStorage::Streaming
1534 );
1535
1536 assert_eq!(
1537 selected.map_events(|event| event.scalar(0)).unwrap(),
1538 vec![0.0, 1.0, 2.0]
1539 );
1540 assert_eq!(reads.load(Ordering::Relaxed), 2);
1541 assert_eq!(
1542 bins[0]
1543 .dataset()
1544 .map_events(|event| event.scalar(0))
1545 .unwrap(),
1546 vec![0.0]
1547 );
1548 assert_eq!(reads.load(Ordering::Relaxed), 2);
1549 }
1550
1551 #[test]
1552 fn unknown_cardinality_fastest_discovers_and_retains_small_selection() {
1553 let source = dataset();
1554 let batch = source.batches().unwrap().next().unwrap().unwrap();
1555 let reads = Arc::new(AtomicUsize::new(0));
1556 let dataset = Dataset::new(CountingSource {
1557 inner: MemorySource::new(batch),
1558 reads: Arc::clone(&reads),
1559 });
1560 let execution = Execution::default();
1561 let x = event_scalar("x");
1562 let selected = dataset
1563 .select(&Predicate::ge(x.clone(), 0.0), &execution)
1564 .unwrap();
1565 let compiled = CompiledModel::from_expr(&x).unwrap();
1566 let params = compiled.params().default_values();
1567 let model = PreparedModel::prepare(&compiled, &execution).unwrap();
1568 let prepared = model.prepare_dataset(&execution, &selected).unwrap();
1569
1570 #[cfg(not(feature = "wgpu"))]
1571 let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared;
1572 #[cfg(feature = "wgpu")]
1573 let crate::PreparedDataset::Cpu(prepared_cpu) = &prepared else {
1574 panic!("default execution prepares CPU datasets");
1575 };
1576 assert_eq!(
1577 prepared_cpu.stats().storage(),
1578 laddu_data::data::CacheStorage::Resident
1579 );
1580 assert_eq!(prepared_cpu.stats().local_events(), 3);
1581 assert_eq!(reads.load(Ordering::Relaxed), 2);
1582
1583 for _ in 0..2 {
1584 assert_eq!(
1585 model
1586 .reduce(
1587 &execution,
1588 ¶ms,
1589 &prepared,
1590 laddu_compile::ReductionPlan::weighted_real(),
1591 )
1592 .unwrap(),
1593 5.5
1594 );
1595 }
1596 assert_eq!(reads.load(Ordering::Relaxed), 2);
1597 }
1598}