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