1use std::sync::{
2 Arc, Mutex,
3 atomic::{AtomicU64, Ordering},
4};
5
6use crate::{
7 LadduDataError, LadduDataResult,
8 data::event::{Event, EventBatch, OwnedEvent},
9 io::{EventSink, EventSource, ReadPlan, SourceCapabilities, WritePlan, memory::MemorySource},
10 schema::Schema,
11};
12use laddu_memory::{MemoryBudget, MemoryDecision};
13use num::complex::Complex64;
14
15#[cfg(feature = "parallel")]
16pub mod accurate;
17mod execution;
18mod ops;
19
20use execution::{DatasetExecutionPlan, DatasetExecutor, visit_events};
21use ops::DatasetOp;
22#[cfg(test)]
23use ops::{poisson1_from_hash, uniform_hash_01};
24
25static NEXT_DATASET_IDENTITY: AtomicU64 = AtomicU64::new(1);
26
27fn next_dataset_identity() -> u64 {
28 NEXT_DATASET_IDENTITY.fetch_add(1, Ordering::Relaxed)
29}
30
31#[derive(Copy, Clone, Debug, PartialEq)]
33pub struct DatasetStats {
34 events: u64,
35 sum_weights: f64,
36 sum_squared_weights: f64,
37 positive_weights: f64,
38 negative_weights: f64,
39}
40
41impl DatasetStats {
42 pub fn events(&self) -> u64 {
44 self.events
45 }
46
47 pub fn sum_weights(&self) -> f64 {
49 self.sum_weights
50 }
51
52 pub fn sum_squared_weights(&self) -> f64 {
54 self.sum_squared_weights
55 }
56
57 pub fn effective_entries(&self) -> Option<f64> {
62 (self.sum_squared_weights > 0.0)
63 .then(|| self.sum_weights * self.sum_weights / self.sum_squared_weights)
64 }
65
66 pub fn positive_weights(&self) -> f64 {
68 self.positive_weights
69 }
70
71 pub fn negative_weights(&self) -> f64 {
73 self.negative_weights
74 }
75}
76
77#[derive(Default)]
78struct DatasetStatsCache {
79 events: Option<u64>,
80 sum_weights: Option<f64>,
81 sum_squared_weights: Option<f64>,
82 positive_weights: Option<f64>,
83 negative_weights: Option<f64>,
84}
85
86#[derive(Clone)]
88pub struct Dataset {
89 identity: u64,
90 source: Arc<dyn EventSource>,
91 plan: ReadPlan,
92 ops: Arc<[DatasetOp]>,
93 cache_storage: CacheStorage,
94 memory_policy: MemoryPolicy,
95 memory_budget: MemoryBudget,
96 last_memory_decision: Arc<Mutex<Option<MemoryDecision>>>,
97 stats: Arc<Mutex<DatasetStatsCache>>,
98 source_traversals: Arc<AtomicU64>,
99}
100
101#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
103pub enum CacheStorage {
104 #[default]
106 Resident,
107 Streaming,
109}
110
111#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
113pub enum MemoryPolicy {
114 #[default]
116 Fastest,
117 Resident,
119 Streaming,
121}
122
123impl Dataset {
124 pub fn new<S>(source: S) -> Self
126 where
127 S: EventSource + 'static,
128 {
129 Self {
130 identity: next_dataset_identity(),
131 source: Arc::new(source),
132 plan: ReadPlan::default(),
133 ops: Arc::from([]),
134 cache_storage: CacheStorage::Resident,
135 memory_policy: MemoryPolicy::Fastest,
136 memory_budget: MemoryBudget::Auto,
137 last_memory_decision: Default::default(),
138 stats: Default::default(),
139 source_traversals: Default::default(),
140 }
141 }
142
143 pub fn from_arc(source: Arc<dyn EventSource>) -> Self {
145 Self {
146 identity: next_dataset_identity(),
147 source,
148 plan: ReadPlan::default(),
149 ops: Arc::from([]),
150 cache_storage: CacheStorage::Resident,
151 memory_policy: MemoryPolicy::Fastest,
152 memory_budget: MemoryBudget::Auto,
153 last_memory_decision: Default::default(),
154 stats: Default::default(),
155 source_traversals: Default::default(),
156 }
157 }
158
159 pub fn with_derived_source<S>(&self, source: S) -> Self
164 where
165 S: EventSource + 'static,
166 {
167 Self {
168 identity: next_dataset_identity(),
169 source: Arc::new(source),
170 plan: self.plan,
171 ops: Arc::from([]),
172 cache_storage: self.cache_storage,
173 memory_policy: self.memory_policy,
174 memory_budget: self.memory_budget,
175 last_memory_decision: Default::default(),
176 stats: Default::default(),
177 source_traversals: Default::default(),
178 }
179 }
180
181 pub fn empty_derived(&self) -> LadduDataResult<Self> {
187 Ok(self.with_derived_source(MemorySource::empty(self.schema()?)))
188 }
189
190 pub fn from_batch(batch: EventBatch) -> Self {
192 Self::new(MemorySource::new(batch))
193 }
194
195 pub fn from_batches(batches: Vec<EventBatch>) -> LadduDataResult<Self> {
201 Ok(Self::new(MemorySource::from_batches(batches)?))
202 }
203
204 pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
210 where
211 I: IntoIterator<Item = OwnedEvent>,
212 {
213 Ok(Self::new(MemorySource::from_events(schema, events)?))
214 }
215
216 pub fn schema(&self) -> LadduDataResult<Arc<Schema>> {
223 self.source.schema()
224 }
225
226 pub fn column(&self, name: &str) -> LadduDataResult<crate::columns::Column> {
231 let schema = self.schema()?;
232 let index = schema
233 .column_index(name)
234 .ok_or_else(|| LadduDataError::MissingColumn(crate::Name::from(name)))?;
235 let mut values = crate::columns::ColumnBuffer::new(schema.columns()[index].1, 0);
236 for batch in self.batches()? {
237 values.extend(batch?.column(index))?;
238 }
239 Ok(values.finish())
240 }
241
242 pub fn capabilities(&self) -> SourceCapabilities {
244 self.source.capabilities()
245 }
246
247 pub fn num_events(&self) -> LadduDataResult<Option<u64>> {
253 {
254 let stats = self.stats.lock().unwrap_or_else(|error| error.into_inner());
255 if let Some(events) = stats.events {
256 return Ok(Some(events));
257 }
258 }
259
260 if self
261 .ops
262 .iter()
263 .any(|op| !matches!(op, DatasetOp::Bootstrap { .. }))
264 {
265 return Ok(None);
266 }
267
268 let events = self.source.num_events()?;
269 if let Some(events) = events {
270 self.stats
271 .lock()
272 .unwrap_or_else(|error| error.into_inner())
273 .events = Some(events);
274 }
275 Ok(events)
276 }
277
278 pub fn stats(&self) -> LadduDataResult<DatasetStats> {
287 {
288 let cache = self.stats.lock().unwrap_or_else(|error| error.into_inner());
289 if let (
290 Some(events),
291 Some(sum_weights),
292 Some(sum_squared_weights),
293 Some(positive_weights),
294 Some(negative_weights),
295 ) = (
296 cache.events,
297 cache.sum_weights,
298 cache.sum_squared_weights,
299 cache.positive_weights,
300 cache.negative_weights,
301 ) {
302 return Ok(DatasetStats {
303 events,
304 sum_weights,
305 sum_squared_weights,
306 positive_weights,
307 negative_weights,
308 });
309 }
310 }
311
312 let mut executor = self.executor_with_plan(self.plan)?;
313 for batch in &mut executor {
314 batch?;
315 }
316
317 Ok(executor.stats())
318 }
319
320 pub fn read_plan(&self) -> ReadPlan {
322 self.plan
323 }
324
325 pub fn cache_storage(&self) -> CacheStorage {
327 self.cache_storage
328 }
329
330 pub fn memory_policy(&self) -> MemoryPolicy {
332 self.memory_policy
333 }
334
335 pub fn memory_budget(&self) -> MemoryBudget {
337 self.memory_budget
338 }
339
340 pub fn last_memory_decision(&self) -> Option<MemoryDecision> {
342 self.last_memory_decision
343 .lock()
344 .unwrap_or_else(|error| error.into_inner())
345 .clone()
346 }
347
348 pub fn source_traversals(&self) -> u64 {
350 self.source_traversals.load(Ordering::Relaxed)
351 }
352
353 #[doc(hidden)]
358 pub fn identity(&self) -> u64 {
359 self.identity
360 }
361
362 pub fn with_memory_budget(mut self, budget: MemoryBudget) -> Self {
364 self.memory_budget = budget;
365 self
366 }
367
368 pub fn fastest(mut self) -> Self {
370 self.memory_policy = MemoryPolicy::Fastest;
371 self.cache_storage = CacheStorage::Resident;
372 self
373 }
374
375 pub fn resident(mut self) -> Self {
381 self.memory_policy = MemoryPolicy::Resident;
382 self.cache_storage = CacheStorage::Resident;
383 self
384 }
385
386 pub fn streaming(mut self) -> Self {
391 self.memory_policy = MemoryPolicy::Streaming;
392 self.cache_storage = CacheStorage::Streaming;
393 self
394 }
395
396 pub fn chunked(mut self, chunk_size: usize) -> LadduDataResult<Self> {
406 if chunk_size == 0 {
407 return Err(LadduDataError::InvalidArgument(
408 "chunk_size must be nonzero",
409 ));
410 }
411 self.plan.chunk_size = Some(chunk_size);
412 Ok(self)
413 }
414
415 pub fn unchunked(mut self) -> Self {
417 self.plan.chunk_size = None;
418 self
419 }
420
421 pub fn filter<F>(self, f: F) -> Self
423 where
424 F: Fn(Event<'_>) -> bool + Send + Sync + 'static,
425 {
426 self.push_op(DatasetOp::Filter(Arc::new(f)))
427 }
428
429 pub fn subsample(self, fraction: f64, seed: u64) -> LadduDataResult<Self> {
436 if !(0.0..=1.0).contains(&fraction) {
437 return Err(LadduDataError::InvalidArgument(
438 "fraction must be in [0, 1]",
439 ));
440 }
441
442 Ok(self.push_op(DatasetOp::Subsample { fraction, seed }))
443 }
444
445 pub fn bootstrap(self, seed: u64) -> Self {
447 self.push_op(DatasetOp::Bootstrap { seed })
448 }
449
450 pub fn for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
457 where
458 F: FnMut(Event<'_>),
459 {
460 self.try_for_each_event(|ev| {
461 f(ev);
462 Ok(())
463 })
464 }
465
466 pub fn try_for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
473 where
474 F: FnMut(Event<'_>) -> LadduDataResult<()>,
475 {
476 visit_events(
477 self,
478 DatasetExecutionPlan::resolve(self, self.plan)?,
479 &mut f,
480 )
481 }
482
483 pub fn try_map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
490 where
491 F: FnMut(Event<'_>) -> LadduDataResult<T>,
492 {
493 let mut out = Vec::new();
494 self.try_for_each_event(|ev| {
495 out.push(f(ev)?);
496 Ok(())
497 })?;
498
499 Ok(out)
500 }
501
502 pub fn map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
509 where
510 F: FnMut(Event<'_>) -> T,
511 {
512 let mut out = Vec::new();
513 self.try_for_each_event(|ev| {
514 out.push(f(ev));
515 Ok(())
516 })?;
517
518 Ok(out)
519 }
520
521 pub fn try_fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
528 where
529 F: FnMut(T, Event<'_>) -> LadduDataResult<T>,
530 {
531 let mut acc = Some(init);
532
533 self.try_for_each_event(|ev| {
534 let current = acc.take().ok_or_else(|| {
535 LadduDataError::Source("dataset fold accumulator was consumed".into())
536 })?;
537 acc = Some(f(current, ev)?);
538 Ok(())
539 })?;
540
541 acc.ok_or_else(|| LadduDataError::Source("dataset fold produced no accumulator".into()))
542 }
543
544 pub fn fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
551 where
552 F: FnMut(T, Event<'_>) -> T,
553 {
554 self.try_fold_events(init, |acc, ev| Ok(f(acc, ev)))
555 }
556
557 pub fn try_fold_batches<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
566 where
567 F: FnMut(T, EventBatch) -> LadduDataResult<T>,
568 {
569 let mut acc = init;
570 for batch in self.batches()? {
571 acc = f(acc, batch?)?;
572 }
573 Ok(acc)
574 }
575
576 pub fn try_accumulate_events<T, F>(&self, mut acc: T, mut f: F) -> LadduDataResult<T>
583 where
584 F: FnMut(&mut T, Event<'_>) -> LadduDataResult<()>,
585 {
586 self.try_for_each_event(|ev| f(&mut acc, ev))?;
587 Ok(acc)
588 }
589
590 pub fn accumulate_events<T, F>(&self, acc: T, mut f: F) -> LadduDataResult<T>
597 where
598 F: FnMut(&mut T, Event<'_>),
599 {
600 self.try_accumulate_events(acc, |acc, ev| {
601 f(acc, ev);
602 Ok(())
603 })
604 }
605
606 pub fn sum_weights(&self) -> LadduDataResult<f64> {
613 Ok(self.stats()?.sum_weights())
614 }
615
616 pub fn weighted_sum<F>(&self, mut f: F) -> LadduDataResult<f64>
623 where
624 F: FnMut(Event<'_>) -> f64,
625 {
626 self.fold_events(0.0, |sum, ev| sum + ev.weight() * f(ev))
627 }
628
629 pub fn weighted_complex_sum<F>(&self, mut f: F) -> LadduDataResult<Complex64>
636 where
637 F: FnMut(Event<'_>) -> Complex64,
638 {
639 self.fold_events(0.0.into(), |sum, ev| sum + ev.weight() * f(ev))
640 }
641
642 pub fn batches(
649 &self,
650 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
651 self.stream_with_plan(self.plan)
652 }
653
654 #[doc(hidden)]
655 pub fn stream_with_plan(
662 &self,
663 plan: ReadPlan,
664 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
665 Ok(Box::new(self.executor_with_plan(plan)?))
666 }
667
668 #[doc(hidden)]
669 pub fn batches_with_plan(
671 &self,
672 plan: ReadPlan,
673 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
674 self.stream_with_plan(plan)
675 }
676
677 fn executor_with_plan(&self, plan: ReadPlan) -> LadduDataResult<DatasetExecutor> {
678 DatasetExecutor::new(self, DatasetExecutionPlan::resolve(self, plan)?)
679 }
680
681 pub fn try_for_each_batch<F>(&self, mut f: F) -> LadduDataResult<()>
688 where
689 F: FnMut(EventBatch) -> LadduDataResult<()>,
690 {
691 for batch in self.batches()? {
692 f(batch?)?;
693 }
694
695 Ok(())
696 }
697
698 pub fn map_batches<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
704 where
705 F: FnMut(EventBatch) -> T,
706 {
707 let mut out = Vec::new();
708
709 self.try_for_each_batch(|batch| {
710 out.push(f(batch));
711 Ok(())
712 })?;
713
714 Ok(out)
715 }
716
717 pub fn write_to<S: EventSink>(&self, sink: &mut S) -> LadduDataResult<()> {
724 sink.begin(self.schema()?, WritePlan::from(self.plan))?;
725
726 let result = (|| {
727 for batch in self.batches()? {
728 sink.write_batch(&batch?)?;
729 }
730
731 sink.finish()
732 })();
733
734 if result.is_err() {
735 let _ = sink.abort();
738 }
739
740 result
741 }
742
743 fn push_op(self, op: DatasetOp) -> Self {
744 let preserved_events = if matches!(&op, DatasetOp::Bootstrap { .. }) {
745 self.num_events().ok().flatten()
746 } else {
747 None
748 };
749 let mut ops = self.ops.to_vec();
750 ops.push(op);
751
752 Self {
753 identity: next_dataset_identity(),
754 source: self.source,
755 plan: self.plan,
756 ops: ops.into(),
757 cache_storage: self.cache_storage,
758 memory_policy: self.memory_policy,
759 memory_budget: self.memory_budget,
760 last_memory_decision: Default::default(),
761 stats: Arc::new(Mutex::new(DatasetStatsCache {
762 events: preserved_events,
763 sum_weights: None,
764 sum_squared_weights: None,
765 positive_weights: None,
766 negative_weights: None,
767 })),
768 source_traversals: Default::default(),
769 }
770 }
771}
772
773#[cfg(test)]
774mod tests {
775 use super::ops::materialize_batch;
776 use super::*;
777 use crate::io::{EventBatchIter, EventSource, ReadPlan, memory::MemorySink};
778 use laddu_physics::vectors::RealVec4;
779 use std::sync::atomic::{AtomicUsize, Ordering};
780
781 #[derive(Clone)]
782 struct CountingSource {
783 batch: EventBatch,
784 reads: Arc<AtomicUsize>,
785 }
786
787 impl EventSource for CountingSource {
788 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
789 Ok(Arc::clone(self.batch.schema()))
790 }
791
792 fn batches(
793 &self,
794 _plan: ReadPlan,
795 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
796 self.reads.fetch_add(1, Ordering::Relaxed);
797 Ok(Box::new(std::iter::once(Ok(self.batch.clone()))))
798 }
799 }
800
801 #[derive(Clone)]
802 struct ErrorSource {
803 schema: Arc<Schema>,
804 items: Arc<[LadduDataResult<EventBatch>]>,
805 }
806
807 impl EventSource for ErrorSource {
808 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
809 Ok(Arc::clone(&self.schema))
810 }
811
812 fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
813 let items = Arc::clone(&self.items);
814 Ok(Box::new(
815 (0..items.len()).map(move |index| items[index].clone()),
816 ))
817 }
818 }
819
820 fn v(x: f64) -> RealVec4 {
821 RealVec4 {
822 e: x + 0.3,
823 px: x,
824 py: x + 0.1,
825 pz: x + 0.2,
826 }
827 }
828
829 fn schema_with_weight() -> Arc<Schema> {
830 Arc::new(Schema::new(["p"], ["x"], true).unwrap())
831 }
832
833 fn schema_without_weight() -> Arc<Schema> {
834 Arc::new(Schema::new(["p"], ["x"], false).unwrap())
835 }
836
837 fn weighted_batch(start: usize, len: usize) -> EventBatch {
838 let schema = schema_with_weight();
839
840 let events = (start..start + len)
841 .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
842
843 EventBatch::from_events(schema, events).unwrap()
844 }
845
846 fn unweighted_batch(start: usize, len: usize) -> EventBatch {
847 let schema = schema_without_weight();
848
849 let events =
850 (start..start + len).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64]));
851
852 EventBatch::from_events(schema, events).unwrap()
853 }
854
855 fn error_source(error_after: Option<EventBatch>) -> ErrorSource {
856 let schema = schema_with_weight();
857 let mut items = Vec::new();
858 if let Some(batch) = error_after {
859 items.push(Ok(batch));
860 }
861 items.push(Err(LadduDataError::Unsupported("source")));
862 ErrorSource {
863 schema,
864 items: items.into(),
865 }
866 }
867
868 fn scalar_values(batch: &EventBatch) -> Vec<f64> {
869 batch.scalar_column(0).to_vec()
870 }
871
872 fn stats_dataset(weights: &[f64]) -> Dataset {
873 let schema = schema_with_weight();
874 let events = weights.iter().enumerate().map(|(index, &weight)| {
875 OwnedEvent::weighted(vec![v(index as f64)], vec![index as f64], weight)
876 });
877 Dataset::from_batch(EventBatch::from_events(schema, events).unwrap())
878 }
879
880 #[test]
881 fn dataset_statistics_are_shared_per_view_and_invalidated_by_selection() {
882 let reads = Arc::new(AtomicUsize::new(0));
883 let dataset = Dataset::new(CountingSource {
884 batch: weighted_batch(0, 5),
885 reads: Arc::clone(&reads),
886 });
887 let clone = dataset.clone();
888
889 assert_eq!(dataset.num_events().unwrap(), None);
890 assert_eq!(dataset.stats().unwrap().events(), 5);
891 assert_eq!(clone.sum_weights().unwrap(), 60.0);
892 let stats = clone.stats().unwrap();
893 assert_eq!(stats.sum_squared_weights(), 730.0);
894 assert_eq!(stats.positive_weights(), 60.0);
895 assert_eq!(stats.negative_weights(), 0.0);
896 assert_eq!(stats.effective_entries(), Some(3600.0 / 730.0));
897 assert_eq!(clone.num_events().unwrap(), Some(5));
898 assert_eq!(reads.load(Ordering::Relaxed), 1);
899
900 let bootstrapped = dataset.clone().bootstrap(7);
901 assert_eq!(bootstrapped.num_events().unwrap(), Some(5));
902 assert_eq!(reads.load(Ordering::Relaxed), 1);
903
904 let filtered = dataset.filter(|event| event.scalar(0) >= 2.0);
905 assert_eq!(filtered.num_events().unwrap(), None);
906 assert_eq!(
907 filtered.map_events(|event| event.scalar(0)).unwrap(),
908 [2.0, 3.0, 4.0]
909 );
910 assert_eq!(reads.load(Ordering::Relaxed), 2);
911 assert_eq!(filtered.stats().unwrap().events(), 3);
912 assert_eq!(reads.load(Ordering::Relaxed), 2);
913 }
914
915 #[test]
916 fn dataset_statistics_preserve_signed_and_squared_weight_diagnostics() {
917 let stats = stats_dataset(&[2.0, -1.0, 3.0, -4.0]).stats().unwrap();
918
919 assert_eq!(stats.events(), 4);
920 assert_eq!(stats.sum_weights(), 0.0);
921 assert_eq!(stats.sum_squared_weights(), 30.0);
922 assert_eq!(stats.positive_weights(), 5.0);
923 assert_eq!(stats.negative_weights(), -5.0);
924 assert_eq!(stats.effective_entries(), Some(0.0));
925 }
926
927 #[test]
928 fn dataset_statistics_make_effective_entries_unavailable_without_squared_weight() {
929 let zero = stats_dataset(&[0.0, 0.0]).stats().unwrap();
930 assert_eq!(zero.events(), 2);
931 assert_eq!(zero.sum_weights(), 0.0);
932 assert_eq!(zero.sum_squared_weights(), 0.0);
933 assert_eq!(zero.positive_weights(), 0.0);
934 assert_eq!(zero.negative_weights(), 0.0);
935 assert_eq!(zero.effective_entries(), None);
936
937 let empty = stats_dataset(&[1.0])
938 .empty_derived()
939 .unwrap()
940 .stats()
941 .unwrap();
942 assert_eq!(empty.events(), 0);
943 assert_eq!(empty.sum_weights(), 0.0);
944 assert_eq!(empty.sum_squared_weights(), 0.0);
945 assert_eq!(empty.effective_entries(), None);
946 }
947
948 #[test]
949 fn dataset_statistics_match_across_resident_streaming_and_chunked_views() {
950 let resident = Dataset::from_batch(weighted_batch(0, 10));
951 let streaming = Dataset::new(CountingSource {
952 batch: weighted_batch(0, 10),
953 reads: Arc::new(AtomicUsize::new(0)),
954 });
955 let chunked = Dataset::from_batches(
956 (0..10)
957 .map(|index| weighted_batch(index, 1))
958 .collect::<Vec<_>>(),
959 )
960 .unwrap()
961 .chunked(3)
962 .unwrap();
963
964 let expected = resident.stats().unwrap();
965 assert_eq!(streaming.stats().unwrap(), expected);
966 assert_eq!(chunked.stats().unwrap(), expected);
967 }
968
969 #[test]
970 fn transformed_fragments_are_coalesced_to_the_read_chunk_size() {
971 let fragments = (0..10)
972 .map(|index| weighted_batch(index, 1))
973 .collect::<Vec<_>>();
974 let dataset = Dataset::from_batches(fragments)
975 .unwrap()
976 .filter(|_| true)
977 .chunked(4)
978 .unwrap();
979
980 let batches = dataset
981 .batches()
982 .unwrap()
983 .collect::<LadduDataResult<Vec<_>>>()
984 .unwrap();
985 assert_eq!(
986 batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
987 [4, 4, 2]
988 );
989 assert_eq!(
990 batches
991 .iter()
992 .flat_map(|batch| batch.scalar_column(0).iter().copied())
993 .collect::<Vec<_>>(),
994 (0..10).map(|value| value as f64).collect::<Vec<_>>()
995 );
996 }
997
998 #[test]
999 fn shared_stream_preserves_pending_batches_before_source_errors() {
1000 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
1001 .chunked(4)
1002 .unwrap();
1003 let mut batches = dataset.batches().unwrap();
1004
1005 assert_eq!(batches.next().unwrap().unwrap().len(), 2);
1006 assert!(matches!(
1007 batches.next().unwrap(),
1008 Err(LadduDataError::Unsupported("source"))
1009 ));
1010 assert!(batches.next().is_none());
1011
1012 assert!(matches!(
1013 dataset.stats(),
1014 Err(LadduDataError::Unsupported("source"))
1015 ));
1016 assert_eq!(dataset.source_traversals(), 2);
1017 }
1018
1019 #[test]
1020 fn event_visitors_share_the_execution_plan_without_changing_source_rows() {
1021 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
1022 .chunked(4)
1023 .unwrap()
1024 .filter(|event| event.scalar(0) >= 0.0);
1025 let mut rows = Vec::new();
1026
1027 let error = dataset
1028 .try_for_each_event(|event| {
1029 rows.push(event.row());
1030 Ok(())
1031 })
1032 .unwrap_err();
1033
1034 assert!(matches!(error, LadduDataError::Unsupported("source")));
1035 assert_eq!(rows, [0, 1]);
1036 }
1037
1038 #[test]
1039 fn dataset_map_fold_accumulate_complex_sum_and_error_paths_use_transformed_events() {
1040 let dataset =
1041 Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) % 2.0 == 0.0);
1042
1043 let rows = dataset
1044 .map_events(|ev| (ev.row(), ev.scalar(0), ev.weight()))
1045 .unwrap();
1046
1047 assert_eq!(rows, vec![(0, 0.0, 10.0), (2, 2.0, 12.0), (4, 4.0, 14.0)]);
1048
1049 let folded = dataset
1050 .fold_events(String::new(), |mut out, ev| {
1051 out.push_str(&format!("{};", ev.scalar(0)));
1052 out
1053 })
1054 .unwrap();
1055
1056 assert_eq!(folded, "0;2;4;");
1057
1058 let accumulated = dataset
1059 .accumulate_events(Vec::<f64>::new(), |values, ev| values.push(ev.weight()))
1060 .unwrap();
1061
1062 assert_eq!(accumulated, vec![10.0, 12.0, 14.0]);
1063
1064 let weighted_sum = dataset.weighted_sum(|ev| ev.scalar(0)).unwrap();
1065 assert_eq!(weighted_sum, 0.0 * 10.0 + 2.0 * 12.0 + 4.0 * 14.0);
1066
1067 let complex_sum = dataset
1068 .weighted_complex_sum(|ev| Complex64::new(ev.scalar(0), 1.0))
1069 .unwrap();
1070
1071 assert_eq!(complex_sum.re, weighted_sum);
1072 assert_eq!(complex_sum.im, 10.0 + 12.0 + 14.0);
1073
1074 let err = dataset
1075 .try_map_events(|ev| {
1076 if ev.scalar(0) == 2.0 {
1077 Err(LadduDataError::Unsupported("stop"))
1078 } else {
1079 Ok(ev.scalar(0))
1080 }
1081 })
1082 .unwrap_err();
1083
1084 assert!(matches!(err, LadduDataError::Unsupported("stop")));
1085 }
1086
1087 #[test]
1088 fn batch_folds_and_empty_derived_sources_preserve_schema_and_errors() {
1089 let dataset = Dataset::from_batches(vec![weighted_batch(0, 2), weighted_batch(2, 2)])
1090 .unwrap()
1091 .chunked(2)
1092 .unwrap();
1093 let event_count = dataset
1094 .try_fold_batches(0usize, |count, batch| Ok(count + batch.len()))
1095 .unwrap();
1096 assert_eq!(event_count, 4);
1097 let rows = dataset
1098 .try_fold_batches(Vec::new(), |mut rows, batch| {
1099 rows.extend((0..batch.len()).map(|row| batch.scalar_at(0, row)));
1100 Ok(rows)
1101 })
1102 .unwrap();
1103 assert_eq!(rows, [0.0, 1.0, 2.0, 3.0]);
1104
1105 let error = dataset
1106 .try_fold_batches(0usize, |_count, _batch| {
1107 Err(LadduDataError::Unsupported("stop"))
1108 })
1109 .unwrap_err();
1110 assert!(matches!(error, LadduDataError::Unsupported("stop")));
1111
1112 let empty = dataset.empty_derived().unwrap();
1113 assert_eq!(
1114 empty.schema().unwrap().as_ref(),
1115 dataset.schema().unwrap().as_ref()
1116 );
1117 assert_eq!(empty.num_events().unwrap(), Some(0));
1118 assert!(empty.batches().unwrap().next().is_none());
1119 }
1120
1121 #[test]
1122 fn deterministic_subsample_and_bootstrap_use_global_event_ids_across_batches() {
1123 let seed = 0x0BAD_5EED;
1124 let bootstrap_seed = 0xB007_57A9;
1125
1126 let dataset = Dataset::from_batches(vec![weighted_batch(0, 3), weighted_batch(3, 3)])
1127 .unwrap()
1128 .subsample(0.5, seed)
1129 .unwrap()
1130 .bootstrap(bootstrap_seed);
1131
1132 let observed = dataset
1133 .map_events(|ev| (ev.scalar(0) as u64, ev.weight()))
1134 .unwrap();
1135
1136 let expected: Vec<(u64, f64)> = (0_u64..6)
1137 .filter(|&event_id| uniform_hash_01(seed, event_id) < 0.5)
1138 .map(|event_id| {
1139 let original_weight = 10.0 + event_id as f64;
1140 let bootstrap_weight =
1141 poisson1_from_hash(bootstrap_seed, event_id) as f64 * original_weight;
1142 (event_id, bootstrap_weight)
1143 })
1144 .collect();
1145
1146 assert_eq!(observed, expected);
1147 }
1148
1149 #[test]
1150 fn materialized_batches_store_weights_only_when_needed() {
1151 let unweighted = unweighted_batch(0, 4);
1152
1153 let filtered = Dataset::from_batch(unweighted.clone())
1154 .filter(|ev| ev.scalar(0) >= 1.0)
1155 .subsample(1.0, 123)
1156 .unwrap();
1157
1158 let filtered_batch = filtered.batches().unwrap().next().unwrap().unwrap();
1159
1160 assert_eq!(scalar_values(&filtered_batch), vec![1.0, 2.0, 3.0]);
1161 assert!(filtered_batch.weights_column().is_none());
1162
1163 let bootstrapped = Dataset::from_batch(unweighted).bootstrap(999);
1164 let bootstrapped_batch = bootstrapped.batches().unwrap().next().unwrap().unwrap();
1165
1166 assert!(bootstrapped_batch.weights_column().is_some());
1167
1168 let source = weighted_batch(0, 2);
1169 let empty_weighted =
1170 materialize_batch(&source, &[DatasetOp::Filter(Arc::new(|_| false))], 0).unwrap();
1171 assert!(empty_weighted.is_empty());
1172 assert_eq!(empty_weighted.weights_column(), Some([].as_slice()));
1173 }
1174
1175 #[test]
1176 fn write_to_memory_sink_captures_transformed_dataset() {
1177 let dataset = Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) >= 2.0);
1178
1179 let mut sink = MemorySink::new();
1180 dataset.write_to(&mut sink).unwrap();
1181
1182 let captured = sink.into_batch().unwrap();
1183
1184 assert_eq!(scalar_values(&captured), vec![2.0, 3.0, 4.0]);
1185 assert_eq!(captured.weights_column().unwrap(), &[12.0, 13.0, 14.0]);
1186 }
1187
1188 #[test]
1189 fn immutable_dataset_identity_tracks_semantic_views() {
1190 let dataset = Dataset::from_batch(weighted_batch(0, 3));
1191 assert_eq!(dataset.identity(), dataset.clone().identity());
1192 assert_eq!(dataset.identity(), dataset.clone().streaming().identity());
1193 assert_ne!(
1194 dataset.identity(),
1195 dataset.clone().subsample(1.0, 7).unwrap().identity()
1196 );
1197 assert_ne!(dataset.identity(), dataset.clone().bootstrap(7).identity());
1198 assert_ne!(
1199 dataset.identity(),
1200 dataset.clone().filter(|_| true).identity()
1201 );
1202 }
1203}