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 capabilities(&self) -> SourceCapabilities {
228 self.source.capabilities()
229 }
230
231 pub fn num_events(&self) -> LadduDataResult<Option<u64>> {
237 {
238 let stats = self.stats.lock().unwrap_or_else(|error| error.into_inner());
239 if let Some(events) = stats.events {
240 return Ok(Some(events));
241 }
242 }
243
244 if self
245 .ops
246 .iter()
247 .any(|op| !matches!(op, DatasetOp::Bootstrap { .. }))
248 {
249 return Ok(None);
250 }
251
252 let events = self.source.num_events()?;
253 if let Some(events) = events {
254 self.stats
255 .lock()
256 .unwrap_or_else(|error| error.into_inner())
257 .events = Some(events);
258 }
259 Ok(events)
260 }
261
262 pub fn stats(&self) -> LadduDataResult<DatasetStats> {
271 {
272 let cache = self.stats.lock().unwrap_or_else(|error| error.into_inner());
273 if let (
274 Some(events),
275 Some(sum_weights),
276 Some(sum_squared_weights),
277 Some(positive_weights),
278 Some(negative_weights),
279 ) = (
280 cache.events,
281 cache.sum_weights,
282 cache.sum_squared_weights,
283 cache.positive_weights,
284 cache.negative_weights,
285 ) {
286 return Ok(DatasetStats {
287 events,
288 sum_weights,
289 sum_squared_weights,
290 positive_weights,
291 negative_weights,
292 });
293 }
294 }
295
296 let mut executor = self.executor_with_plan(self.plan)?;
297 for batch in &mut executor {
298 batch?;
299 }
300
301 Ok(executor.stats())
302 }
303
304 pub fn read_plan(&self) -> ReadPlan {
306 self.plan
307 }
308
309 pub fn cache_storage(&self) -> CacheStorage {
311 self.cache_storage
312 }
313
314 pub fn memory_policy(&self) -> MemoryPolicy {
316 self.memory_policy
317 }
318
319 pub fn memory_budget(&self) -> MemoryBudget {
321 self.memory_budget
322 }
323
324 pub fn last_memory_decision(&self) -> Option<MemoryDecision> {
326 self.last_memory_decision
327 .lock()
328 .unwrap_or_else(|error| error.into_inner())
329 .clone()
330 }
331
332 pub fn source_traversals(&self) -> u64 {
334 self.source_traversals.load(Ordering::Relaxed)
335 }
336
337 #[doc(hidden)]
342 pub fn identity(&self) -> u64 {
343 self.identity
344 }
345
346 pub fn with_memory_budget(mut self, budget: MemoryBudget) -> Self {
348 self.memory_budget = budget;
349 self
350 }
351
352 pub fn fastest(mut self) -> Self {
354 self.memory_policy = MemoryPolicy::Fastest;
355 self.cache_storage = CacheStorage::Resident;
356 self
357 }
358
359 pub fn resident(mut self) -> Self {
365 self.memory_policy = MemoryPolicy::Resident;
366 self.cache_storage = CacheStorage::Resident;
367 self
368 }
369
370 pub fn streaming(mut self) -> Self {
375 self.memory_policy = MemoryPolicy::Streaming;
376 self.cache_storage = CacheStorage::Streaming;
377 self
378 }
379
380 pub fn chunked(mut self, chunk_size: usize) -> LadduDataResult<Self> {
390 if chunk_size == 0 {
391 return Err(LadduDataError::InvalidArgument(
392 "chunk_size must be nonzero",
393 ));
394 }
395 self.plan.chunk_size = Some(chunk_size);
396 Ok(self)
397 }
398
399 pub fn unchunked(mut self) -> Self {
401 self.plan.chunk_size = None;
402 self
403 }
404
405 pub fn filter<F>(self, f: F) -> Self
407 where
408 F: Fn(Event<'_>) -> bool + Send + Sync + 'static,
409 {
410 self.push_op(DatasetOp::Filter(Arc::new(f)))
411 }
412
413 pub fn subsample(self, fraction: f64, seed: u64) -> LadduDataResult<Self> {
420 if !(0.0..=1.0).contains(&fraction) {
421 return Err(LadduDataError::InvalidArgument(
422 "fraction must be in [0, 1]",
423 ));
424 }
425
426 Ok(self.push_op(DatasetOp::Subsample { fraction, seed }))
427 }
428
429 pub fn bootstrap(self, seed: u64) -> Self {
431 self.push_op(DatasetOp::Bootstrap { seed })
432 }
433
434 pub fn for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
441 where
442 F: FnMut(Event<'_>),
443 {
444 self.try_for_each_event(|ev| {
445 f(ev);
446 Ok(())
447 })
448 }
449
450 pub fn try_for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
457 where
458 F: FnMut(Event<'_>) -> LadduDataResult<()>,
459 {
460 visit_events(
461 self,
462 DatasetExecutionPlan::resolve(self, self.plan)?,
463 &mut f,
464 )
465 }
466
467 pub fn try_map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
474 where
475 F: FnMut(Event<'_>) -> LadduDataResult<T>,
476 {
477 let mut out = Vec::new();
478 self.try_for_each_event(|ev| {
479 out.push(f(ev)?);
480 Ok(())
481 })?;
482
483 Ok(out)
484 }
485
486 pub fn map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
493 where
494 F: FnMut(Event<'_>) -> T,
495 {
496 let mut out = Vec::new();
497 self.try_for_each_event(|ev| {
498 out.push(f(ev));
499 Ok(())
500 })?;
501
502 Ok(out)
503 }
504
505 pub fn try_fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
512 where
513 F: FnMut(T, Event<'_>) -> LadduDataResult<T>,
514 {
515 let mut acc = Some(init);
516
517 self.try_for_each_event(|ev| {
518 let current = acc.take().ok_or_else(|| {
519 LadduDataError::Source("dataset fold accumulator was consumed".into())
520 })?;
521 acc = Some(f(current, ev)?);
522 Ok(())
523 })?;
524
525 acc.ok_or_else(|| LadduDataError::Source("dataset fold produced no accumulator".into()))
526 }
527
528 pub fn fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
535 where
536 F: FnMut(T, Event<'_>) -> T,
537 {
538 self.try_fold_events(init, |acc, ev| Ok(f(acc, ev)))
539 }
540
541 pub fn try_fold_batches<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
550 where
551 F: FnMut(T, EventBatch) -> LadduDataResult<T>,
552 {
553 let mut acc = init;
554 for batch in self.batches()? {
555 acc = f(acc, batch?)?;
556 }
557 Ok(acc)
558 }
559
560 pub fn try_accumulate_events<T, F>(&self, mut acc: T, mut f: F) -> LadduDataResult<T>
567 where
568 F: FnMut(&mut T, Event<'_>) -> LadduDataResult<()>,
569 {
570 self.try_for_each_event(|ev| f(&mut acc, ev))?;
571 Ok(acc)
572 }
573
574 pub fn accumulate_events<T, F>(&self, acc: T, mut f: F) -> LadduDataResult<T>
581 where
582 F: FnMut(&mut T, Event<'_>),
583 {
584 self.try_accumulate_events(acc, |acc, ev| {
585 f(acc, ev);
586 Ok(())
587 })
588 }
589
590 pub fn sum_weights(&self) -> LadduDataResult<f64> {
597 Ok(self.stats()?.sum_weights())
598 }
599
600 pub fn weighted_sum<F>(&self, mut f: F) -> LadduDataResult<f64>
607 where
608 F: FnMut(Event<'_>) -> f64,
609 {
610 self.fold_events(0.0, |sum, ev| sum + ev.weight() * f(ev))
611 }
612
613 pub fn weighted_complex_sum<F>(&self, mut f: F) -> LadduDataResult<Complex64>
620 where
621 F: FnMut(Event<'_>) -> Complex64,
622 {
623 self.fold_events(0.0.into(), |sum, ev| sum + ev.weight() * f(ev))
624 }
625
626 pub fn batches(
633 &self,
634 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
635 self.stream_with_plan(self.plan)
636 }
637
638 #[doc(hidden)]
639 pub fn stream_with_plan(
646 &self,
647 plan: ReadPlan,
648 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
649 Ok(Box::new(self.executor_with_plan(plan)?))
650 }
651
652 #[doc(hidden)]
653 pub fn batches_with_plan(
655 &self,
656 plan: ReadPlan,
657 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
658 self.stream_with_plan(plan)
659 }
660
661 fn executor_with_plan(&self, plan: ReadPlan) -> LadduDataResult<DatasetExecutor> {
662 DatasetExecutor::new(self, DatasetExecutionPlan::resolve(self, plan)?)
663 }
664
665 pub fn try_for_each_batch<F>(&self, mut f: F) -> LadduDataResult<()>
672 where
673 F: FnMut(EventBatch) -> LadduDataResult<()>,
674 {
675 for batch in self.batches()? {
676 f(batch?)?;
677 }
678
679 Ok(())
680 }
681
682 pub fn map_batches<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
688 where
689 F: FnMut(EventBatch) -> T,
690 {
691 let mut out = Vec::new();
692
693 self.try_for_each_batch(|batch| {
694 out.push(f(batch));
695 Ok(())
696 })?;
697
698 Ok(out)
699 }
700
701 pub fn write_to<S: EventSink>(&self, sink: &mut S) -> LadduDataResult<()> {
708 sink.begin(self.schema()?, WritePlan::from(self.plan))?;
709
710 let result = (|| {
711 for batch in self.batches()? {
712 sink.write_batch(&batch?)?;
713 }
714
715 sink.finish()
716 })();
717
718 if result.is_err() {
719 let _ = sink.abort();
722 }
723
724 result
725 }
726
727 fn push_op(self, op: DatasetOp) -> Self {
728 let preserved_events = if matches!(&op, DatasetOp::Bootstrap { .. }) {
729 self.num_events().ok().flatten()
730 } else {
731 None
732 };
733 let mut ops = self.ops.to_vec();
734 ops.push(op);
735
736 Self {
737 identity: next_dataset_identity(),
738 source: self.source,
739 plan: self.plan,
740 ops: ops.into(),
741 cache_storage: self.cache_storage,
742 memory_policy: self.memory_policy,
743 memory_budget: self.memory_budget,
744 last_memory_decision: Default::default(),
745 stats: Arc::new(Mutex::new(DatasetStatsCache {
746 events: preserved_events,
747 sum_weights: None,
748 sum_squared_weights: None,
749 positive_weights: None,
750 negative_weights: None,
751 })),
752 source_traversals: Default::default(),
753 }
754 }
755}
756
757#[cfg(test)]
758mod tests {
759 use super::ops::materialize_batch;
760 use super::*;
761 use crate::io::{EventBatchIter, EventSource, ReadPlan, memory::MemorySink};
762 use laddu_physics::vectors::RealVec4;
763 use std::sync::atomic::{AtomicUsize, Ordering};
764
765 #[derive(Clone)]
766 struct CountingSource {
767 batch: EventBatch,
768 reads: Arc<AtomicUsize>,
769 }
770
771 impl EventSource for CountingSource {
772 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
773 Ok(Arc::clone(self.batch.schema()))
774 }
775
776 fn batches(
777 &self,
778 _plan: ReadPlan,
779 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
780 self.reads.fetch_add(1, Ordering::Relaxed);
781 Ok(Box::new(std::iter::once(Ok(self.batch.clone()))))
782 }
783 }
784
785 #[derive(Clone)]
786 struct ErrorSource {
787 schema: Arc<Schema>,
788 items: Arc<[LadduDataResult<EventBatch>]>,
789 }
790
791 impl EventSource for ErrorSource {
792 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
793 Ok(Arc::clone(&self.schema))
794 }
795
796 fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
797 let items = Arc::clone(&self.items);
798 Ok(Box::new(
799 (0..items.len()).map(move |index| items[index].clone()),
800 ))
801 }
802 }
803
804 fn v(x: f64) -> RealVec4 {
805 RealVec4 {
806 e: x + 0.3,
807 px: x,
808 py: x + 0.1,
809 pz: x + 0.2,
810 }
811 }
812
813 fn schema_with_weight() -> Arc<Schema> {
814 Arc::new(Schema::new(["p"], ["x"], true).unwrap())
815 }
816
817 fn schema_without_weight() -> Arc<Schema> {
818 Arc::new(Schema::new(["p"], ["x"], false).unwrap())
819 }
820
821 fn weighted_batch(start: usize, len: usize) -> EventBatch {
822 let schema = schema_with_weight();
823
824 let events = (start..start + len)
825 .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
826
827 EventBatch::from_events(schema, events).unwrap()
828 }
829
830 fn unweighted_batch(start: usize, len: usize) -> EventBatch {
831 let schema = schema_without_weight();
832
833 let events =
834 (start..start + len).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64]));
835
836 EventBatch::from_events(schema, events).unwrap()
837 }
838
839 fn error_source(error_after: Option<EventBatch>) -> ErrorSource {
840 let schema = schema_with_weight();
841 let mut items = Vec::new();
842 if let Some(batch) = error_after {
843 items.push(Ok(batch));
844 }
845 items.push(Err(LadduDataError::Unsupported("source")));
846 ErrorSource {
847 schema,
848 items: items.into(),
849 }
850 }
851
852 fn scalar_values(batch: &EventBatch) -> Vec<f64> {
853 batch.scalar_column(0).to_vec()
854 }
855
856 fn stats_dataset(weights: &[f64]) -> Dataset {
857 let schema = schema_with_weight();
858 let events = weights.iter().enumerate().map(|(index, &weight)| {
859 OwnedEvent::weighted(vec![v(index as f64)], vec![index as f64], weight)
860 });
861 Dataset::from_batch(EventBatch::from_events(schema, events).unwrap())
862 }
863
864 #[test]
865 fn dataset_statistics_are_shared_per_view_and_invalidated_by_selection() {
866 let reads = Arc::new(AtomicUsize::new(0));
867 let dataset = Dataset::new(CountingSource {
868 batch: weighted_batch(0, 5),
869 reads: Arc::clone(&reads),
870 });
871 let clone = dataset.clone();
872
873 assert_eq!(dataset.num_events().unwrap(), None);
874 assert_eq!(dataset.stats().unwrap().events(), 5);
875 assert_eq!(clone.sum_weights().unwrap(), 60.0);
876 let stats = clone.stats().unwrap();
877 assert_eq!(stats.sum_squared_weights(), 730.0);
878 assert_eq!(stats.positive_weights(), 60.0);
879 assert_eq!(stats.negative_weights(), 0.0);
880 assert_eq!(stats.effective_entries(), Some(3600.0 / 730.0));
881 assert_eq!(clone.num_events().unwrap(), Some(5));
882 assert_eq!(reads.load(Ordering::Relaxed), 1);
883
884 let bootstrapped = dataset.clone().bootstrap(7);
885 assert_eq!(bootstrapped.num_events().unwrap(), Some(5));
886 assert_eq!(reads.load(Ordering::Relaxed), 1);
887
888 let filtered = dataset.filter(|event| event.scalar(0) >= 2.0);
889 assert_eq!(filtered.num_events().unwrap(), None);
890 assert_eq!(
891 filtered.map_events(|event| event.scalar(0)).unwrap(),
892 [2.0, 3.0, 4.0]
893 );
894 assert_eq!(reads.load(Ordering::Relaxed), 2);
895 assert_eq!(filtered.stats().unwrap().events(), 3);
896 assert_eq!(reads.load(Ordering::Relaxed), 2);
897 }
898
899 #[test]
900 fn dataset_statistics_preserve_signed_and_squared_weight_diagnostics() {
901 let stats = stats_dataset(&[2.0, -1.0, 3.0, -4.0]).stats().unwrap();
902
903 assert_eq!(stats.events(), 4);
904 assert_eq!(stats.sum_weights(), 0.0);
905 assert_eq!(stats.sum_squared_weights(), 30.0);
906 assert_eq!(stats.positive_weights(), 5.0);
907 assert_eq!(stats.negative_weights(), -5.0);
908 assert_eq!(stats.effective_entries(), Some(0.0));
909 }
910
911 #[test]
912 fn dataset_statistics_make_effective_entries_unavailable_without_squared_weight() {
913 let zero = stats_dataset(&[0.0, 0.0]).stats().unwrap();
914 assert_eq!(zero.events(), 2);
915 assert_eq!(zero.sum_weights(), 0.0);
916 assert_eq!(zero.sum_squared_weights(), 0.0);
917 assert_eq!(zero.positive_weights(), 0.0);
918 assert_eq!(zero.negative_weights(), 0.0);
919 assert_eq!(zero.effective_entries(), None);
920
921 let empty = stats_dataset(&[1.0])
922 .empty_derived()
923 .unwrap()
924 .stats()
925 .unwrap();
926 assert_eq!(empty.events(), 0);
927 assert_eq!(empty.sum_weights(), 0.0);
928 assert_eq!(empty.sum_squared_weights(), 0.0);
929 assert_eq!(empty.effective_entries(), None);
930 }
931
932 #[test]
933 fn dataset_statistics_match_across_resident_streaming_and_chunked_views() {
934 let resident = Dataset::from_batch(weighted_batch(0, 10));
935 let streaming = Dataset::new(CountingSource {
936 batch: weighted_batch(0, 10),
937 reads: Arc::new(AtomicUsize::new(0)),
938 });
939 let chunked = Dataset::from_batches(
940 (0..10)
941 .map(|index| weighted_batch(index, 1))
942 .collect::<Vec<_>>(),
943 )
944 .unwrap()
945 .chunked(3)
946 .unwrap();
947
948 let expected = resident.stats().unwrap();
949 assert_eq!(streaming.stats().unwrap(), expected);
950 assert_eq!(chunked.stats().unwrap(), expected);
951 }
952
953 #[test]
954 fn transformed_fragments_are_coalesced_to_the_read_chunk_size() {
955 let fragments = (0..10)
956 .map(|index| weighted_batch(index, 1))
957 .collect::<Vec<_>>();
958 let dataset = Dataset::from_batches(fragments)
959 .unwrap()
960 .filter(|_| true)
961 .chunked(4)
962 .unwrap();
963
964 let batches = dataset
965 .batches()
966 .unwrap()
967 .collect::<LadduDataResult<Vec<_>>>()
968 .unwrap();
969 assert_eq!(
970 batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
971 [4, 4, 2]
972 );
973 assert_eq!(
974 batches
975 .iter()
976 .flat_map(|batch| batch.scalar_column(0).iter().copied())
977 .collect::<Vec<_>>(),
978 (0..10).map(|value| value as f64).collect::<Vec<_>>()
979 );
980 }
981
982 #[test]
983 fn shared_stream_preserves_pending_batches_before_source_errors() {
984 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
985 .chunked(4)
986 .unwrap();
987 let mut batches = dataset.batches().unwrap();
988
989 assert_eq!(batches.next().unwrap().unwrap().len(), 2);
990 assert!(matches!(
991 batches.next().unwrap(),
992 Err(LadduDataError::Unsupported("source"))
993 ));
994 assert!(batches.next().is_none());
995
996 assert!(matches!(
997 dataset.stats(),
998 Err(LadduDataError::Unsupported("source"))
999 ));
1000 assert_eq!(dataset.source_traversals(), 2);
1001 }
1002
1003 #[test]
1004 fn event_visitors_share_the_execution_plan_without_changing_source_rows() {
1005 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
1006 .chunked(4)
1007 .unwrap()
1008 .filter(|event| event.scalar(0) >= 0.0);
1009 let mut rows = Vec::new();
1010
1011 let error = dataset
1012 .try_for_each_event(|event| {
1013 rows.push(event.row());
1014 Ok(())
1015 })
1016 .unwrap_err();
1017
1018 assert!(matches!(error, LadduDataError::Unsupported("source")));
1019 assert_eq!(rows, [0, 1]);
1020 }
1021
1022 #[test]
1023 fn dataset_map_fold_accumulate_complex_sum_and_error_paths_use_transformed_events() {
1024 let dataset =
1025 Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) % 2.0 == 0.0);
1026
1027 let rows = dataset
1028 .map_events(|ev| (ev.row(), ev.scalar(0), ev.weight()))
1029 .unwrap();
1030
1031 assert_eq!(rows, vec![(0, 0.0, 10.0), (2, 2.0, 12.0), (4, 4.0, 14.0)]);
1032
1033 let folded = dataset
1034 .fold_events(String::new(), |mut out, ev| {
1035 out.push_str(&format!("{};", ev.scalar(0)));
1036 out
1037 })
1038 .unwrap();
1039
1040 assert_eq!(folded, "0;2;4;");
1041
1042 let accumulated = dataset
1043 .accumulate_events(Vec::<f64>::new(), |values, ev| values.push(ev.weight()))
1044 .unwrap();
1045
1046 assert_eq!(accumulated, vec![10.0, 12.0, 14.0]);
1047
1048 let weighted_sum = dataset.weighted_sum(|ev| ev.scalar(0)).unwrap();
1049 assert_eq!(weighted_sum, 0.0 * 10.0 + 2.0 * 12.0 + 4.0 * 14.0);
1050
1051 let complex_sum = dataset
1052 .weighted_complex_sum(|ev| Complex64::new(ev.scalar(0), 1.0))
1053 .unwrap();
1054
1055 assert_eq!(complex_sum.re, weighted_sum);
1056 assert_eq!(complex_sum.im, 10.0 + 12.0 + 14.0);
1057
1058 let err = dataset
1059 .try_map_events(|ev| {
1060 if ev.scalar(0) == 2.0 {
1061 Err(LadduDataError::Unsupported("stop"))
1062 } else {
1063 Ok(ev.scalar(0))
1064 }
1065 })
1066 .unwrap_err();
1067
1068 assert!(matches!(err, LadduDataError::Unsupported("stop")));
1069 }
1070
1071 #[test]
1072 fn batch_folds_and_empty_derived_sources_preserve_schema_and_errors() {
1073 let dataset = Dataset::from_batches(vec![weighted_batch(0, 2), weighted_batch(2, 2)])
1074 .unwrap()
1075 .chunked(2)
1076 .unwrap();
1077 let event_count = dataset
1078 .try_fold_batches(0usize, |count, batch| Ok(count + batch.len()))
1079 .unwrap();
1080 assert_eq!(event_count, 4);
1081 let rows = dataset
1082 .try_fold_batches(Vec::new(), |mut rows, batch| {
1083 rows.extend((0..batch.len()).map(|row| batch.scalar_at(0, row)));
1084 Ok(rows)
1085 })
1086 .unwrap();
1087 assert_eq!(rows, [0.0, 1.0, 2.0, 3.0]);
1088
1089 let error = dataset
1090 .try_fold_batches(0usize, |_count, _batch| {
1091 Err(LadduDataError::Unsupported("stop"))
1092 })
1093 .unwrap_err();
1094 assert!(matches!(error, LadduDataError::Unsupported("stop")));
1095
1096 let empty = dataset.empty_derived().unwrap();
1097 assert_eq!(
1098 empty.schema().unwrap().as_ref(),
1099 dataset.schema().unwrap().as_ref()
1100 );
1101 assert_eq!(empty.num_events().unwrap(), Some(0));
1102 assert!(empty.batches().unwrap().next().is_none());
1103 }
1104
1105 #[test]
1106 fn deterministic_subsample_and_bootstrap_use_global_event_ids_across_batches() {
1107 let seed = 0x0BAD_5EED;
1108 let bootstrap_seed = 0xB007_57A9;
1109
1110 let dataset = Dataset::from_batches(vec![weighted_batch(0, 3), weighted_batch(3, 3)])
1111 .unwrap()
1112 .subsample(0.5, seed)
1113 .unwrap()
1114 .bootstrap(bootstrap_seed);
1115
1116 let observed = dataset
1117 .map_events(|ev| (ev.scalar(0) as u64, ev.weight()))
1118 .unwrap();
1119
1120 let expected: Vec<(u64, f64)> = (0_u64..6)
1121 .filter(|&event_id| uniform_hash_01(seed, event_id) < 0.5)
1122 .map(|event_id| {
1123 let original_weight = 10.0 + event_id as f64;
1124 let bootstrap_weight =
1125 poisson1_from_hash(bootstrap_seed, event_id) as f64 * original_weight;
1126 (event_id, bootstrap_weight)
1127 })
1128 .collect();
1129
1130 assert_eq!(observed, expected);
1131 }
1132
1133 #[test]
1134 fn materialized_batches_store_weights_only_when_needed() {
1135 let unweighted = unweighted_batch(0, 4);
1136
1137 let filtered = Dataset::from_batch(unweighted.clone())
1138 .filter(|ev| ev.scalar(0) >= 1.0)
1139 .subsample(1.0, 123)
1140 .unwrap();
1141
1142 let filtered_batch = filtered.batches().unwrap().next().unwrap().unwrap();
1143
1144 assert_eq!(scalar_values(&filtered_batch), vec![1.0, 2.0, 3.0]);
1145 assert!(filtered_batch.weights_column().is_none());
1146
1147 let bootstrapped = Dataset::from_batch(unweighted).bootstrap(999);
1148 let bootstrapped_batch = bootstrapped.batches().unwrap().next().unwrap().unwrap();
1149
1150 assert!(bootstrapped_batch.weights_column().is_some());
1151
1152 let source = weighted_batch(0, 2);
1153 let empty_weighted =
1154 materialize_batch(&source, &[DatasetOp::Filter(Arc::new(|_| false))], 0).unwrap();
1155 assert!(empty_weighted.is_empty());
1156 assert_eq!(empty_weighted.weights_column(), Some([].as_slice()));
1157 }
1158
1159 #[test]
1160 fn write_to_memory_sink_captures_transformed_dataset() {
1161 let dataset = Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) >= 2.0);
1162
1163 let mut sink = MemorySink::new();
1164 dataset.write_to(&mut sink).unwrap();
1165
1166 let captured = sink.into_batch().unwrap();
1167
1168 assert_eq!(scalar_values(&captured), vec![2.0, 3.0, 4.0]);
1169 assert_eq!(captured.weights_column().unwrap(), &[12.0, 13.0, 14.0]);
1170 }
1171
1172 #[test]
1173 fn immutable_dataset_identity_tracks_semantic_views() {
1174 let dataset = Dataset::from_batch(weighted_batch(0, 3));
1175 assert_eq!(dataset.identity(), dataset.clone().identity());
1176 assert_eq!(dataset.identity(), dataset.clone().streaming().identity());
1177 assert_ne!(
1178 dataset.identity(),
1179 dataset.clone().subsample(1.0, 7).unwrap().identity()
1180 );
1181 assert_ne!(dataset.identity(), dataset.clone().bootstrap(7).identity());
1182 assert_ne!(
1183 dataset.identity(),
1184 dataset.clone().filter(|_| true).identity()
1185 );
1186 }
1187}