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}
37
38impl DatasetStats {
39 pub fn events(&self) -> u64 {
41 self.events
42 }
43
44 pub fn sum_weights(&self) -> f64 {
46 self.sum_weights
47 }
48}
49
50#[derive(Default)]
51struct DatasetStatsCache {
52 events: Option<u64>,
53 sum_weights: Option<f64>,
54}
55
56#[derive(Clone)]
58pub struct Dataset {
59 identity: u64,
60 source: Arc<dyn EventSource>,
61 plan: ReadPlan,
62 ops: Arc<[DatasetOp]>,
63 cache_storage: CacheStorage,
64 memory_policy: MemoryPolicy,
65 memory_budget: MemoryBudget,
66 last_memory_decision: Arc<Mutex<Option<MemoryDecision>>>,
67 stats: Arc<Mutex<DatasetStatsCache>>,
68 source_traversals: Arc<AtomicU64>,
69}
70
71#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
73pub enum CacheStorage {
74 #[default]
76 Resident,
77 Streaming,
79}
80
81#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
83pub enum MemoryPolicy {
84 #[default]
86 Fastest,
87 Resident,
89 Streaming,
91}
92
93impl Dataset {
94 pub fn new<S>(source: S) -> Self
96 where
97 S: EventSource + 'static,
98 {
99 Self {
100 identity: next_dataset_identity(),
101 source: Arc::new(source),
102 plan: ReadPlan::default(),
103 ops: Arc::from([]),
104 cache_storage: CacheStorage::Resident,
105 memory_policy: MemoryPolicy::Fastest,
106 memory_budget: MemoryBudget::Auto,
107 last_memory_decision: Default::default(),
108 stats: Default::default(),
109 source_traversals: Default::default(),
110 }
111 }
112
113 pub fn from_arc(source: Arc<dyn EventSource>) -> Self {
115 Self {
116 identity: next_dataset_identity(),
117 source,
118 plan: ReadPlan::default(),
119 ops: Arc::from([]),
120 cache_storage: CacheStorage::Resident,
121 memory_policy: MemoryPolicy::Fastest,
122 memory_budget: MemoryBudget::Auto,
123 last_memory_decision: Default::default(),
124 stats: Default::default(),
125 source_traversals: Default::default(),
126 }
127 }
128
129 #[doc(hidden)]
131 pub fn with_derived_source<S>(&self, source: S) -> Self
132 where
133 S: EventSource + 'static,
134 {
135 Self {
136 identity: next_dataset_identity(),
137 source: Arc::new(source),
138 plan: self.plan,
139 ops: Arc::from([]),
140 cache_storage: self.cache_storage,
141 memory_policy: self.memory_policy,
142 memory_budget: self.memory_budget,
143 last_memory_decision: Default::default(),
144 stats: Default::default(),
145 source_traversals: Default::default(),
146 }
147 }
148
149 pub fn from_batch(batch: EventBatch) -> Self {
151 Self::new(MemorySource::new(batch))
152 }
153
154 pub fn from_batches(batches: Vec<EventBatch>) -> LadduDataResult<Self> {
160 Ok(Self::new(MemorySource::from_batches(batches)?))
161 }
162
163 pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
169 where
170 I: IntoIterator<Item = OwnedEvent>,
171 {
172 Ok(Self::new(MemorySource::from_events(schema, events)?))
173 }
174
175 pub fn schema(&self) -> LadduDataResult<Arc<Schema>> {
182 self.source.schema()
183 }
184
185 pub fn capabilities(&self) -> SourceCapabilities {
187 self.source.capabilities()
188 }
189
190 pub fn num_events(&self) -> LadduDataResult<Option<u64>> {
196 {
197 let stats = self.stats.lock().unwrap_or_else(|error| error.into_inner());
198 if let Some(events) = stats.events {
199 return Ok(Some(events));
200 }
201 }
202
203 if self
204 .ops
205 .iter()
206 .any(|op| !matches!(op, DatasetOp::Bootstrap { .. }))
207 {
208 return Ok(None);
209 }
210
211 let events = self.source.num_events()?;
212 if let Some(events) = events {
213 self.stats
214 .lock()
215 .unwrap_or_else(|error| error.into_inner())
216 .events = Some(events);
217 }
218 Ok(events)
219 }
220
221 pub fn stats(&self) -> LadduDataResult<DatasetStats> {
230 {
231 let cache = self.stats.lock().unwrap_or_else(|error| error.into_inner());
232 if let (Some(events), Some(sum_weights)) = (cache.events, cache.sum_weights) {
233 return Ok(DatasetStats {
234 events,
235 sum_weights,
236 });
237 }
238 }
239
240 if self.ops.is_empty()
241 && let (Some(events), Some(sum_weights)) =
242 (self.source.num_events()?, self.source.weighted_total()?)
243 {
244 let stats = DatasetStats {
245 events,
246 sum_weights,
247 };
248 let mut cache = self.stats.lock().unwrap_or_else(|error| error.into_inner());
249 cache.events = Some(events);
250 cache.sum_weights = Some(sum_weights);
251 return Ok(stats);
252 }
253
254 let mut executor = self.executor_with_plan(self.plan)?;
255 for batch in &mut executor {
256 batch?;
257 }
258
259 Ok(executor.stats())
260 }
261
262 pub fn read_plan(&self) -> ReadPlan {
264 self.plan
265 }
266
267 pub fn cache_storage(&self) -> CacheStorage {
269 self.cache_storage
270 }
271
272 pub fn memory_policy(&self) -> MemoryPolicy {
274 self.memory_policy
275 }
276
277 pub fn memory_budget(&self) -> MemoryBudget {
279 self.memory_budget
280 }
281
282 pub fn last_memory_decision(&self) -> Option<MemoryDecision> {
284 self.last_memory_decision
285 .lock()
286 .unwrap_or_else(|error| error.into_inner())
287 .clone()
288 }
289
290 pub fn source_traversals(&self) -> u64 {
292 self.source_traversals.load(Ordering::Relaxed)
293 }
294
295 #[doc(hidden)]
300 pub fn identity(&self) -> u64 {
301 self.identity
302 }
303
304 pub fn with_memory_budget(mut self, budget: MemoryBudget) -> Self {
306 self.memory_budget = budget;
307 self
308 }
309
310 pub fn fastest(mut self) -> Self {
312 self.memory_policy = MemoryPolicy::Fastest;
313 self.cache_storage = CacheStorage::Resident;
314 self
315 }
316
317 pub fn resident(mut self) -> Self {
323 self.memory_policy = MemoryPolicy::Resident;
324 self.cache_storage = CacheStorage::Resident;
325 self
326 }
327
328 pub fn streaming(mut self) -> Self {
333 self.memory_policy = MemoryPolicy::Streaming;
334 self.cache_storage = CacheStorage::Streaming;
335 self
336 }
337
338 pub fn chunked(mut self, chunk_size: usize) -> LadduDataResult<Self> {
348 if chunk_size == 0 {
349 return Err(LadduDataError::InvalidArgument(
350 "chunk_size must be nonzero",
351 ));
352 }
353 self.plan.chunk_size = Some(chunk_size);
354 Ok(self)
355 }
356
357 pub fn unchunked(mut self) -> Self {
359 self.plan.chunk_size = None;
360 self
361 }
362
363 pub fn filter<F>(self, f: F) -> Self
365 where
366 F: Fn(Event<'_>) -> bool + Send + Sync + 'static,
367 {
368 self.push_op(DatasetOp::Filter(Arc::new(f)))
369 }
370
371 pub fn subsample(self, fraction: f64, seed: u64) -> LadduDataResult<Self> {
378 if !(0.0..=1.0).contains(&fraction) {
379 return Err(LadduDataError::InvalidArgument(
380 "fraction must be in [0, 1]",
381 ));
382 }
383
384 Ok(self.push_op(DatasetOp::Subsample { fraction, seed }))
385 }
386
387 pub fn bootstrap(self, seed: u64) -> Self {
389 self.push_op(DatasetOp::Bootstrap { seed })
390 }
391
392 pub fn for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
399 where
400 F: FnMut(Event<'_>),
401 {
402 self.try_for_each_event(|ev| {
403 f(ev);
404 Ok(())
405 })
406 }
407
408 pub fn try_for_each_event<F>(&self, mut f: F) -> LadduDataResult<()>
415 where
416 F: FnMut(Event<'_>) -> LadduDataResult<()>,
417 {
418 visit_events(
419 self,
420 DatasetExecutionPlan::resolve(self, self.plan)?,
421 &mut f,
422 )
423 }
424
425 pub fn try_map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
432 where
433 F: FnMut(Event<'_>) -> LadduDataResult<T>,
434 {
435 let mut out = Vec::new();
436 self.try_for_each_event(|ev| {
437 out.push(f(ev)?);
438 Ok(())
439 })?;
440
441 Ok(out)
442 }
443
444 pub fn map_events<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
451 where
452 F: FnMut(Event<'_>) -> T,
453 {
454 let mut out = Vec::new();
455 self.try_for_each_event(|ev| {
456 out.push(f(ev));
457 Ok(())
458 })?;
459
460 Ok(out)
461 }
462
463 pub fn try_fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
470 where
471 F: FnMut(T, Event<'_>) -> LadduDataResult<T>,
472 {
473 let mut acc = Some(init);
474
475 self.try_for_each_event(|ev| {
476 let current = acc.take().ok_or_else(|| {
477 LadduDataError::Source("dataset fold accumulator was consumed".into())
478 })?;
479 acc = Some(f(current, ev)?);
480 Ok(())
481 })?;
482
483 acc.ok_or_else(|| LadduDataError::Source("dataset fold produced no accumulator".into()))
484 }
485
486 pub fn fold_events<T, F>(&self, init: T, mut f: F) -> LadduDataResult<T>
493 where
494 F: FnMut(T, Event<'_>) -> T,
495 {
496 self.try_fold_events(init, |acc, ev| Ok(f(acc, ev)))
497 }
498
499 pub fn try_accumulate_events<T, F>(&self, mut acc: T, mut f: F) -> LadduDataResult<T>
506 where
507 F: FnMut(&mut T, Event<'_>) -> LadduDataResult<()>,
508 {
509 self.try_for_each_event(|ev| f(&mut acc, ev))?;
510 Ok(acc)
511 }
512
513 pub fn accumulate_events<T, F>(&self, acc: T, mut f: F) -> LadduDataResult<T>
520 where
521 F: FnMut(&mut T, Event<'_>),
522 {
523 self.try_accumulate_events(acc, |acc, ev| {
524 f(acc, ev);
525 Ok(())
526 })
527 }
528
529 pub fn sum_weights(&self) -> LadduDataResult<f64> {
536 Ok(self.stats()?.sum_weights())
537 }
538
539 pub fn weighted_sum<F>(&self, mut f: F) -> LadduDataResult<f64>
546 where
547 F: FnMut(Event<'_>) -> f64,
548 {
549 self.fold_events(0.0, |sum, ev| sum + ev.weight() * f(ev))
550 }
551
552 pub fn weighted_complex_sum<F>(&self, mut f: F) -> LadduDataResult<Complex64>
559 where
560 F: FnMut(Event<'_>) -> Complex64,
561 {
562 self.fold_events(0.0.into(), |sum, ev| sum + ev.weight() * f(ev))
563 }
564
565 pub fn batches(
572 &self,
573 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
574 self.stream_with_plan(self.plan)
575 }
576
577 #[doc(hidden)]
578 pub fn stream_with_plan(
585 &self,
586 plan: ReadPlan,
587 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
588 Ok(Box::new(self.executor_with_plan(plan)?))
589 }
590
591 #[doc(hidden)]
592 pub fn batches_with_plan(
594 &self,
595 plan: ReadPlan,
596 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
597 self.stream_with_plan(plan)
598 }
599
600 fn executor_with_plan(&self, plan: ReadPlan) -> LadduDataResult<DatasetExecutor> {
601 DatasetExecutor::new(self, DatasetExecutionPlan::resolve(self, plan)?)
602 }
603
604 pub fn try_for_each_batch<F>(&self, mut f: F) -> LadduDataResult<()>
611 where
612 F: FnMut(EventBatch) -> LadduDataResult<()>,
613 {
614 for batch in self.batches()? {
615 f(batch?)?;
616 }
617
618 Ok(())
619 }
620
621 pub fn map_batches<T, F>(&self, mut f: F) -> LadduDataResult<Vec<T>>
627 where
628 F: FnMut(EventBatch) -> T,
629 {
630 let mut out = Vec::new();
631
632 self.try_for_each_batch(|batch| {
633 out.push(f(batch));
634 Ok(())
635 })?;
636
637 Ok(out)
638 }
639
640 pub fn write_to<S: EventSink>(&self, sink: &mut S) -> LadduDataResult<()> {
647 sink.begin(self.schema()?, WritePlan::from(self.plan))?;
648
649 let result = (|| {
650 for batch in self.batches()? {
651 sink.write_batch(&batch?)?;
652 }
653
654 sink.finish()
655 })();
656
657 if result.is_err() {
658 let _ = sink.abort();
661 }
662
663 result
664 }
665
666 fn push_op(self, op: DatasetOp) -> Self {
667 let preserved_events = if matches!(&op, DatasetOp::Bootstrap { .. }) {
668 self.num_events().ok().flatten()
669 } else {
670 None
671 };
672 let mut ops = self.ops.to_vec();
673 ops.push(op);
674
675 Self {
676 identity: next_dataset_identity(),
677 source: self.source,
678 plan: self.plan,
679 ops: ops.into(),
680 cache_storage: self.cache_storage,
681 memory_policy: self.memory_policy,
682 memory_budget: self.memory_budget,
683 last_memory_decision: Default::default(),
684 stats: Arc::new(Mutex::new(DatasetStatsCache {
685 events: preserved_events,
686 sum_weights: None,
687 })),
688 source_traversals: Default::default(),
689 }
690 }
691}
692
693#[cfg(test)]
694mod tests {
695 use super::ops::materialize_batch;
696 use super::*;
697 use crate::io::{EventBatchIter, EventSource, ReadPlan, memory::MemorySink};
698 use laddu_physics::vectors::RealVec4;
699 use std::sync::atomic::{AtomicUsize, Ordering};
700
701 #[derive(Clone)]
702 struct CountingSource {
703 batch: EventBatch,
704 reads: Arc<AtomicUsize>,
705 }
706
707 impl EventSource for CountingSource {
708 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
709 Ok(Arc::clone(self.batch.schema()))
710 }
711
712 fn batches(
713 &self,
714 _plan: ReadPlan,
715 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
716 self.reads.fetch_add(1, Ordering::Relaxed);
717 Ok(Box::new(std::iter::once(Ok(self.batch.clone()))))
718 }
719 }
720
721 #[derive(Clone)]
722 struct ErrorSource {
723 schema: Arc<Schema>,
724 items: Arc<[LadduDataResult<EventBatch>]>,
725 }
726
727 impl EventSource for ErrorSource {
728 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
729 Ok(Arc::clone(&self.schema))
730 }
731
732 fn batches(&self, _plan: ReadPlan) -> LadduDataResult<EventBatchIter> {
733 let items = Arc::clone(&self.items);
734 Ok(Box::new(
735 (0..items.len()).map(move |index| items[index].clone()),
736 ))
737 }
738 }
739
740 fn v(x: f64) -> RealVec4 {
741 RealVec4 {
742 e: x + 0.3,
743 px: x,
744 py: x + 0.1,
745 pz: x + 0.2,
746 }
747 }
748
749 fn schema_with_weight() -> Arc<Schema> {
750 Arc::new(Schema::new(["p"], ["x"], true).unwrap())
751 }
752
753 fn schema_without_weight() -> Arc<Schema> {
754 Arc::new(Schema::new(["p"], ["x"], false).unwrap())
755 }
756
757 fn weighted_batch(start: usize, len: usize) -> EventBatch {
758 let schema = schema_with_weight();
759
760 let events = (start..start + len)
761 .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
762
763 EventBatch::from_events(schema, events).unwrap()
764 }
765
766 fn unweighted_batch(start: usize, len: usize) -> EventBatch {
767 let schema = schema_without_weight();
768
769 let events =
770 (start..start + len).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64]));
771
772 EventBatch::from_events(schema, events).unwrap()
773 }
774
775 fn error_source(error_after: Option<EventBatch>) -> ErrorSource {
776 let schema = schema_with_weight();
777 let mut items = Vec::new();
778 if let Some(batch) = error_after {
779 items.push(Ok(batch));
780 }
781 items.push(Err(LadduDataError::Unsupported("source")));
782 ErrorSource {
783 schema,
784 items: items.into(),
785 }
786 }
787
788 fn scalar_values(batch: &EventBatch) -> Vec<f64> {
789 batch.scalar_column(0).to_vec()
790 }
791
792 #[test]
793 fn dataset_statistics_are_shared_per_view_and_invalidated_by_selection() {
794 let reads = Arc::new(AtomicUsize::new(0));
795 let dataset = Dataset::new(CountingSource {
796 batch: weighted_batch(0, 5),
797 reads: Arc::clone(&reads),
798 });
799 let clone = dataset.clone();
800
801 assert_eq!(dataset.num_events().unwrap(), None);
802 assert_eq!(dataset.stats().unwrap().events(), 5);
803 assert_eq!(clone.sum_weights().unwrap(), 60.0);
804 assert_eq!(clone.num_events().unwrap(), Some(5));
805 assert_eq!(reads.load(Ordering::Relaxed), 1);
806
807 let bootstrapped = dataset.clone().bootstrap(7);
808 assert_eq!(bootstrapped.num_events().unwrap(), Some(5));
809 assert_eq!(reads.load(Ordering::Relaxed), 1);
810
811 let filtered = dataset.filter(|event| event.scalar(0) >= 2.0);
812 assert_eq!(filtered.num_events().unwrap(), None);
813 assert_eq!(
814 filtered.map_events(|event| event.scalar(0)).unwrap(),
815 [2.0, 3.0, 4.0]
816 );
817 assert_eq!(reads.load(Ordering::Relaxed), 2);
818 assert_eq!(filtered.stats().unwrap().events(), 3);
819 assert_eq!(reads.load(Ordering::Relaxed), 2);
820 }
821
822 #[test]
823 fn transformed_fragments_are_coalesced_to_the_read_chunk_size() {
824 let fragments = (0..10)
825 .map(|index| weighted_batch(index, 1))
826 .collect::<Vec<_>>();
827 let dataset = Dataset::from_batches(fragments)
828 .unwrap()
829 .filter(|_| true)
830 .chunked(4)
831 .unwrap();
832
833 let batches = dataset
834 .batches()
835 .unwrap()
836 .collect::<LadduDataResult<Vec<_>>>()
837 .unwrap();
838 assert_eq!(
839 batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
840 [4, 4, 2]
841 );
842 assert_eq!(
843 batches
844 .iter()
845 .flat_map(|batch| batch.scalar_column(0).iter().copied())
846 .collect::<Vec<_>>(),
847 (0..10).map(|value| value as f64).collect::<Vec<_>>()
848 );
849 }
850
851 #[test]
852 fn shared_stream_preserves_pending_batches_before_source_errors() {
853 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
854 .chunked(4)
855 .unwrap();
856 let mut batches = dataset.batches().unwrap();
857
858 assert_eq!(batches.next().unwrap().unwrap().len(), 2);
859 assert!(matches!(
860 batches.next().unwrap(),
861 Err(LadduDataError::Unsupported("source"))
862 ));
863 assert!(batches.next().is_none());
864
865 assert!(matches!(
866 dataset.stats(),
867 Err(LadduDataError::Unsupported("source"))
868 ));
869 assert_eq!(dataset.source_traversals(), 2);
870 }
871
872 #[test]
873 fn event_visitors_share_the_execution_plan_without_changing_source_rows() {
874 let dataset = Dataset::new(error_source(Some(weighted_batch(0, 2))))
875 .chunked(4)
876 .unwrap()
877 .filter(|event| event.scalar(0) >= 0.0);
878 let mut rows = Vec::new();
879
880 let error = dataset
881 .try_for_each_event(|event| {
882 rows.push(event.row());
883 Ok(())
884 })
885 .unwrap_err();
886
887 assert!(matches!(error, LadduDataError::Unsupported("source")));
888 assert_eq!(rows, [0, 1]);
889 }
890
891 #[test]
892 fn dataset_map_fold_accumulate_complex_sum_and_error_paths_use_transformed_events() {
893 let dataset =
894 Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) % 2.0 == 0.0);
895
896 let rows = dataset
897 .map_events(|ev| (ev.row(), ev.scalar(0), ev.weight()))
898 .unwrap();
899
900 assert_eq!(rows, vec![(0, 0.0, 10.0), (2, 2.0, 12.0), (4, 4.0, 14.0)]);
901
902 let folded = dataset
903 .fold_events(String::new(), |mut out, ev| {
904 out.push_str(&format!("{};", ev.scalar(0)));
905 out
906 })
907 .unwrap();
908
909 assert_eq!(folded, "0;2;4;");
910
911 let accumulated = dataset
912 .accumulate_events(Vec::<f64>::new(), |values, ev| values.push(ev.weight()))
913 .unwrap();
914
915 assert_eq!(accumulated, vec![10.0, 12.0, 14.0]);
916
917 let weighted_sum = dataset.weighted_sum(|ev| ev.scalar(0)).unwrap();
918 assert_eq!(weighted_sum, 0.0 * 10.0 + 2.0 * 12.0 + 4.0 * 14.0);
919
920 let complex_sum = dataset
921 .weighted_complex_sum(|ev| Complex64::new(ev.scalar(0), 1.0))
922 .unwrap();
923
924 assert_eq!(complex_sum.re, weighted_sum);
925 assert_eq!(complex_sum.im, 10.0 + 12.0 + 14.0);
926
927 let err = dataset
928 .try_map_events(|ev| {
929 if ev.scalar(0) == 2.0 {
930 Err(LadduDataError::Unsupported("stop"))
931 } else {
932 Ok(ev.scalar(0))
933 }
934 })
935 .unwrap_err();
936
937 assert!(matches!(err, LadduDataError::Unsupported("stop")));
938 }
939
940 #[test]
941 fn deterministic_subsample_and_bootstrap_use_global_event_ids_across_batches() {
942 let seed = 0x0BAD_5EED;
943 let bootstrap_seed = 0xB007_57A9;
944
945 let dataset = Dataset::from_batches(vec![weighted_batch(0, 3), weighted_batch(3, 3)])
946 .unwrap()
947 .subsample(0.5, seed)
948 .unwrap()
949 .bootstrap(bootstrap_seed);
950
951 let observed = dataset
952 .map_events(|ev| (ev.scalar(0) as u64, ev.weight()))
953 .unwrap();
954
955 let expected: Vec<(u64, f64)> = (0_u64..6)
956 .filter(|&event_id| uniform_hash_01(seed, event_id) < 0.5)
957 .map(|event_id| {
958 let original_weight = 10.0 + event_id as f64;
959 let bootstrap_weight =
960 poisson1_from_hash(bootstrap_seed, event_id) as f64 * original_weight;
961 (event_id, bootstrap_weight)
962 })
963 .collect();
964
965 assert_eq!(observed, expected);
966 }
967
968 #[test]
969 fn materialized_batches_store_weights_only_when_needed() {
970 let unweighted = unweighted_batch(0, 4);
971
972 let filtered = Dataset::from_batch(unweighted.clone())
973 .filter(|ev| ev.scalar(0) >= 1.0)
974 .subsample(1.0, 123)
975 .unwrap();
976
977 let filtered_batch = filtered.batches().unwrap().next().unwrap().unwrap();
978
979 assert_eq!(scalar_values(&filtered_batch), vec![1.0, 2.0, 3.0]);
980 assert!(filtered_batch.weights_column().is_none());
981
982 let bootstrapped = Dataset::from_batch(unweighted).bootstrap(999);
983 let bootstrapped_batch = bootstrapped.batches().unwrap().next().unwrap().unwrap();
984
985 assert!(bootstrapped_batch.weights_column().is_some());
986
987 let source = weighted_batch(0, 2);
988 let empty_weighted =
989 materialize_batch(&source, &[DatasetOp::Filter(Arc::new(|_| false))], 0).unwrap();
990 assert!(empty_weighted.is_empty());
991 assert_eq!(empty_weighted.weights_column(), Some([].as_slice()));
992 }
993
994 #[test]
995 fn write_to_memory_sink_captures_transformed_dataset() {
996 let dataset = Dataset::from_batch(weighted_batch(0, 5)).filter(|ev| ev.scalar(0) >= 2.0);
997
998 let mut sink = MemorySink::new();
999 dataset.write_to(&mut sink).unwrap();
1000
1001 let captured = sink.into_batch().unwrap();
1002
1003 assert_eq!(scalar_values(&captured), vec![2.0, 3.0, 4.0]);
1004 assert_eq!(captured.weights_column().unwrap(), &[12.0, 13.0, 14.0]);
1005 }
1006
1007 #[test]
1008 fn immutable_dataset_identity_tracks_semantic_views() {
1009 let dataset = Dataset::from_batch(weighted_batch(0, 3));
1010 assert_eq!(dataset.identity(), dataset.clone().identity());
1011 assert_eq!(dataset.identity(), dataset.clone().streaming().identity());
1012 assert_ne!(
1013 dataset.identity(),
1014 dataset.clone().subsample(1.0, 7).unwrap().identity()
1015 );
1016 assert_ne!(dataset.identity(), dataset.clone().bootstrap(7).identity());
1017 assert_ne!(
1018 dataset.identity(),
1019 dataset.clone().filter(|_| true).identity()
1020 );
1021 }
1022}