1use std::fmt;
2use std::sync::Arc;
3
4use laddu_physics::vectors::RealVec4;
5
6use crate::{
7 BatchLayout, LadduDataError, LadduDataResult,
8 columns::{Column, ColumnBuffer, ColumnValue},
9 schema::{P4Binding, Precision, ScalarBinding, Schema},
10};
11
12#[derive(Clone, Debug)]
13struct BatchParts {
14 p4s: Arc<[Arc<[RealVec4]>]>,
15 scalars: Arc<[Arc<[f64]>]>,
16 columns: Arc<[Column]>,
17 weights: Weights,
18}
19
20#[derive(Clone, Debug)]
21enum Weights {
22 ImplicitUnit,
23 Explicit(Arc<[f64]>),
24}
25
26impl Weights {
27 fn from_option(weights: Option<Arc<[f64]>>) -> Self {
28 match weights {
29 Some(weights) => Self::Explicit(weights),
30 None => Self::ImplicitUnit,
31 }
32 }
33
34 fn as_slice(&self) -> Option<&[f64]> {
35 match self {
36 Self::ImplicitUnit => None,
37 Self::Explicit(weights) => Some(weights),
38 }
39 }
40
41 fn at(&self, row: usize) -> f64 {
42 self.as_slice().map_or(1.0, |weights| weights[row])
43 }
44
45 fn is_explicit(&self) -> bool {
46 matches!(self, Self::Explicit(_))
47 }
48
49 fn select(&self, rows: &[usize]) -> Self {
50 match self {
51 Self::ImplicitUnit => Self::ImplicitUnit,
52 Self::Explicit(weights) => {
53 let selected: Arc<[f64]> = rows.iter().map(|&row| weights[row]).collect();
54 Self::Explicit(selected)
55 }
56 }
57 }
58
59 fn slice(&self, start: usize, end: usize) -> Self {
60 match self {
61 Self::ImplicitUnit => Self::ImplicitUnit,
62 Self::Explicit(weights) => Self::Explicit(Arc::from(&weights[start..end])),
63 }
64 }
65
66 fn reweight<F>(&self, len: usize, f: F) -> Self
67 where
68 F: Fn(usize, f64) -> f64,
69 {
70 let weights: Arc<[f64]> = (0..len).map(|i| f(i, self.at(i))).collect();
71 Self::Explicit(weights)
72 }
73}
74
75impl BatchParts {
76 fn from_columns(p4s: Vec<Arc<[RealVec4]>>, scalars: Vec<Arc<[f64]>>, weights: Weights) -> Self {
77 Self {
78 p4s: p4s.into(),
79 scalars: scalars.into(),
80 columns: Arc::from([]),
81 weights,
82 }
83 }
84
85 fn validate(&self, schema: &Schema, expected_len: Option<usize>) -> LadduDataResult<usize> {
86 if self.p4s.len() != schema.n_p4s() {
87 return Err(LadduDataError::Schema(
88 "wrong number of vec4 columns".into(),
89 ));
90 }
91
92 if self.scalars.len() != schema.n_scalars() {
93 return Err(LadduDataError::Schema(
94 "wrong number of scalar columns".into(),
95 ));
96 }
97
98 if self.columns.len() != schema.n_columns() {
99 return Err(LadduDataError::Schema(
100 "wrong number of typed columns".into(),
101 ));
102 }
103 for (column, (name, dtype)) in self.columns.iter().zip(schema.columns()) {
104 if column.dtype() != *dtype {
105 return Err(LadduDataError::Schema(format!(
106 "column {name:?} dtype does not match schema"
107 )));
108 }
109 }
110 let len = infer_len(
111 &self.p4s,
112 &self.scalars,
113 &self.columns,
114 self.weights.as_slice(),
115 )?;
116 if let Some(expected_len) = expected_len {
117 let has_columns = !self.p4s.is_empty()
118 || !self.scalars.is_empty()
119 || !self.columns.is_empty()
120 || self.weights.is_explicit();
121 if has_columns && len != expected_len {
122 return Err(LadduDataError::Schema("inconsistent batch length".into()));
123 }
124 return Ok(expected_len);
125 }
126 Ok(len)
127 }
128
129 fn select(&self, rows: &[usize]) -> Self {
130 let p4s = self
131 .p4s
132 .iter()
133 .map(|col| rows.iter().map(|&i| col[i]).collect())
134 .collect();
135 let scalars = self
136 .scalars
137 .iter()
138 .map(|col| rows.iter().map(|&i| col[i]).collect())
139 .collect();
140
141 Self {
142 p4s,
143 scalars,
144 weights: self.weights.select(rows),
145 columns: self
146 .columns
147 .iter()
148 .map(|column| column.select(rows))
149 .collect(),
150 }
151 }
152
153 fn slice(&self, start: usize, end: usize) -> Self {
154 let p4s = self
155 .p4s
156 .iter()
157 .map(|col| Arc::<[RealVec4]>::from(&col[start..end]))
158 .collect();
159 let scalars = self
160 .scalars
161 .iter()
162 .map(|col| Arc::<[f64]>::from(&col[start..end]))
163 .collect();
164
165 Self {
166 p4s,
167 scalars,
168 weights: self.weights.slice(start, end),
169 columns: self
170 .columns
171 .iter()
172 .map(|column| column.slice(start, end))
173 .collect(),
174 }
175 }
176
177 fn reweight<F>(&self, len: usize, f: F) -> Self
178 where
179 F: Fn(usize, f64) -> f64,
180 {
181 Self {
182 p4s: Arc::clone(&self.p4s),
183 scalars: Arc::clone(&self.scalars),
184 columns: Arc::clone(&self.columns),
185 weights: self.weights.reweight(len, f),
186 }
187 }
188
189 fn concat(batches: &[(&Self, usize)]) -> LadduDataResult<Self> {
190 let len: usize = batches.iter().map(|(_, len)| *len).sum();
191 let n_p4s = batches.first().map_or(0, |(batch, _)| batch.p4s.len());
192 let n_scalars = batches.first().map_or(0, |(batch, _)| batch.scalars.len());
193
194 let mut p4s = Vec::with_capacity(n_p4s);
195 for col in 0..n_p4s {
196 let mut out = Vec::with_capacity(len);
197 for (batch, _) in batches {
198 out.extend_from_slice(&batch.p4s[col]);
199 }
200 p4s.push(Arc::from(out));
201 }
202
203 let mut scalars = Vec::with_capacity(n_scalars);
204 for col in 0..n_scalars {
205 let mut out = Vec::with_capacity(len);
206 for (batch, _) in batches {
207 out.extend_from_slice(&batch.scalars[col]);
208 }
209 scalars.push(Arc::from(out));
210 }
211
212 let weights = if batches.iter().any(|(batch, _)| batch.weights.is_explicit()) {
213 let mut out = Vec::with_capacity(len);
214 for (batch, batch_len) in batches {
215 for row in 0..*batch_len {
216 out.push(batch.weights.at(row));
217 }
218 }
219 Weights::Explicit(Arc::from(out))
220 } else {
221 Weights::ImplicitUnit
222 };
223
224 let mut parts = Self::from_columns(p4s, scalars, weights);
225 if let Some((first, _)) = batches.first() {
226 parts.columns = first
227 .columns
228 .iter()
229 .enumerate()
230 .map(|(index, column)| {
231 let inputs = batches
232 .iter()
233 .map(|(batch, _)| &batch.columns[index])
234 .collect::<Vec<_>>();
235 Column::concat(column.dtype(), &inputs)
236 })
237 .collect::<LadduDataResult<Vec<_>>>()?
238 .into();
239 }
240 Ok(parts)
241 }
242}
243
244#[derive(Default)]
245enum WeightAssembler {
246 #[default]
247 ImplicitUnit,
248 Explicit(Vec<f64>),
249}
250
251impl WeightAssembler {
252 fn push(&mut self, weight: Option<f64>, len: usize) -> LadduDataResult<()> {
253 match self {
254 Self::Explicit(weights) => match weight {
255 Some(weight) => weights.push(weight),
256 None => {
257 return Err(LadduDataError::InvalidArgument(
258 "cannot mix weighted and unweighted events in one batch",
259 ));
260 }
261 },
262 Self::ImplicitUnit => match weight {
263 Some(weight) if len == 0 => *self = Self::Explicit(vec![weight]),
264 Some(_) => {
265 return Err(LadduDataError::InvalidArgument(
266 "cannot mix unweighted and weighted events in one batch",
267 ));
268 }
269 None => {}
270 },
271 }
272
273 Ok(())
274 }
275
276 fn finish(self) -> Weights {
277 match self {
278 Self::ImplicitUnit => Weights::ImplicitUnit,
279 Self::Explicit(weights) => Weights::Explicit(Arc::from(weights)),
280 }
281 }
282}
283
284#[derive(Clone)]
286pub struct EventBatch {
287 schema: Arc<Schema>,
288 len: usize,
289 parts: BatchParts,
290 row_ids: Option<Arc<[u64]>>,
291}
292
293impl fmt::Debug for EventBatch {
294 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
295 formatter
296 .debug_struct("EventBatch")
297 .field("schema", &self.schema)
298 .field("len", &self.len)
299 .field("p4s", &self.parts.p4s)
300 .field("scalars", &self.parts.scalars)
301 .field("columns", &self.parts.columns)
302 .field("weights", &self.parts.weights.as_slice())
303 .finish()
304 }
305}
306
307impl EventBatch {
308 pub fn new(
315 schema: Arc<Schema>,
316 p4s: Vec<Arc<[RealVec4]>>,
317 scalars: Vec<Arc<[f64]>>,
318 weights: Option<Arc<[f64]>>,
319 ) -> LadduDataResult<Self> {
320 BatchAssembler::from_columns(schema, p4s, scalars, weights)
321 }
322
323 pub fn new_with_columns(
328 schema: Arc<Schema>,
329 p4s: Vec<Arc<[RealVec4]>>,
330 scalars: Vec<Arc<[f64]>>,
331 columns: Vec<Column>,
332 weights: Option<Arc<[f64]>>,
333 ) -> LadduDataResult<Self> {
334 let mut parts = BatchParts::from_columns(p4s, scalars, Weights::from_option(weights));
335 parts.columns = columns.into();
336 Self::from_parts(schema, parts)
337 }
338
339 pub fn new_with_columns_and_len(
344 schema: Arc<Schema>,
345 p4s: Vec<Arc<[RealVec4]>>,
346 scalars: Vec<Arc<[f64]>>,
347 columns: Vec<Column>,
348 weights: Option<Arc<[f64]>>,
349 len: usize,
350 ) -> LadduDataResult<Self> {
351 let mut parts = BatchParts::from_columns(p4s, scalars, Weights::from_option(weights));
352 parts.columns = columns.into();
353 Self::from_parts_with_len(schema, parts, len)
354 }
355
356 pub fn new_with_len(
362 schema: Arc<Schema>,
363 p4s: Vec<Arc<[RealVec4]>>,
364 scalars: Vec<Arc<[f64]>>,
365 weights: Option<Arc<[f64]>>,
366 len: usize,
367 ) -> LadduDataResult<Self> {
368 Self::from_parts_with_len(
369 schema,
370 BatchParts::from_columns(p4s, scalars, Weights::from_option(weights)),
371 len,
372 )
373 }
374
375 fn from_parts(schema: Arc<Schema>, parts: BatchParts) -> LadduDataResult<Self> {
376 let len = parts.validate(&schema, None)?;
377 Ok(Self {
378 schema,
379 len,
380 parts,
381 row_ids: None,
382 })
383 }
384
385 fn from_parts_with_len(
386 schema: Arc<Schema>,
387 parts: BatchParts,
388 expected_len: usize,
389 ) -> LadduDataResult<Self> {
390 let len = parts.validate(&schema, Some(expected_len))?;
391 Ok(Self {
392 schema,
393 len,
394 parts,
395 row_ids: None,
396 })
397 }
398
399 pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
406 where
407 I: IntoIterator<Item = OwnedEvent>,
408 {
409 let mut builder = EventBatchBuilder::new(schema);
410 builder.extend(events)?;
411 builder.finish()
412 }
413
414 pub fn schema(&self) -> &Arc<Schema> {
416 &self.schema
417 }
418
419 pub fn len(&self) -> usize {
421 self.len
422 }
423
424 pub fn row_ids(&self) -> Option<&[u64]> {
426 self.row_ids.as_deref()
427 }
428
429 pub fn with_row_ids(mut self, ids: Arc<[u64]>) -> LadduDataResult<Self> {
434 if ids.len() != self.len {
435 return Err(LadduDataError::Schema(
436 "row identity count does not match event count".into(),
437 ));
438 }
439 self.row_ids = Some(ids);
440 Ok(self)
441 }
442
443 pub fn bytes_per_event(&self) -> usize {
445 BatchLayout::from_batch(self)
446 .bytes_per_event(Precision::F64)
447 .ok()
448 .and_then(|bytes| usize::try_from(bytes).ok())
449 .unwrap_or(usize::MAX)
450 }
451
452 pub fn resident_bytes(&self) -> usize {
457 BatchLayout::from_batch(self)
458 .footprint(Precision::F64)
459 .and_then(|footprint| footprint.checked_peak_bytes(self.len))
460 .ok()
461 .and_then(|bytes| usize::try_from(bytes).ok())
462 .and_then(|bytes| {
463 let identities = self
464 .row_ids
465 .as_ref()
466 .map_or(0, |ids| std::mem::size_of_val(ids.as_ref()));
467 bytes.checked_add(identities)
468 })
469 .unwrap_or(usize::MAX)
470 }
471
472 pub fn is_empty(&self) -> bool {
474 self.len == 0
475 }
476
477 pub fn vec4_column(&self, index: usize) -> &[RealVec4] {
479 &self.parts.p4s[index]
480 }
481
482 pub fn scalar_column(&self, index: usize) -> &[f64] {
484 &self.parts.scalars[index]
485 }
486
487 pub fn column(&self, index: usize) -> &Column {
492 &self.parts.columns[index]
493 }
494
495 pub fn column_named(&self, name: &str) -> Option<&Column> {
497 self.schema
498 .column_index(name)
499 .map(|index| self.column(index))
500 }
501
502 pub fn weights_column(&self) -> Option<&[f64]> {
504 self.parts.weights.as_slice()
505 }
506
507 pub fn vec4_column_named(&self, name: &str) -> Option<&[RealVec4]> {
509 let i = self.schema.p4_index(name)?;
510 Some(self.vec4_column(i))
511 }
512
513 pub fn scalar_column_named(&self, name: &str) -> Option<&[f64]> {
515 let i = self.schema.scalar_index(name)?;
516 Some(self.scalar_column(i))
517 }
518
519 pub fn p4_at(&self, col: usize, row: usize) -> RealVec4 {
521 self.parts.p4s[col][row]
522 }
523
524 pub fn p4_column_bound(&self, binding: &P4Binding) -> LadduDataResult<&[RealVec4]> {
534 if !binding.matches(&self.schema) {
535 return Err(LadduDataError::Schema(
536 "column binding belongs to a different schema".into(),
537 ));
538 }
539 Ok(self.vec4_column(binding.index()))
540 }
541
542 pub fn scalar_at(&self, col: usize, row: usize) -> f64 {
544 self.parts.scalars[col][row]
545 }
546
547 pub fn scalar_column_bound(&self, binding: &ScalarBinding) -> LadduDataResult<&[f64]> {
557 if !binding.matches(&self.schema) {
558 return Err(LadduDataError::Schema(
559 "column binding belongs to a different schema".into(),
560 ));
561 }
562 Ok(self.scalar_column(binding.index()))
563 }
564
565 pub fn weights_at(&self, row: usize) -> f64 {
567 self.parts.weights.at(row)
568 }
569
570 pub fn event(&self, row: usize) -> BatchEvent<'_> {
572 BatchEvent { batch: self, row }
573 }
574
575 pub fn iter(&self) -> impl Iterator<Item = BatchEvent<'_>> {
577 (0..self.len()).map(|i| self.event(i))
578 }
579
580 pub fn select(&self, rows: &[usize]) -> Self {
586 let mut selected = Self::from_parts_with_len(
587 Arc::clone(&self.schema),
588 self.parts.select(rows),
589 rows.len(),
590 )
591 .expect("select preserves EventBatch invariants");
592 selected.row_ids = self
593 .row_ids
594 .as_ref()
595 .map(|ids| rows.iter().map(|&i| ids[i]).collect());
596 selected
597 }
598
599 pub fn filter<F>(&self, keep: F) -> Self
601 where
602 F: Fn(BatchEvent<'_>) -> bool,
603 {
604 let rows: Vec<usize> = (0..self.len).filter(|&i| keep(self.event(i))).collect();
605
606 self.select(&rows)
607 }
608
609 pub fn reweight<F>(&self, f: F) -> Self
616 where
617 F: Fn(usize, f64) -> f64,
618 {
619 let mut reweighted = Self::from_parts_with_len(
620 Arc::clone(&self.schema),
621 self.parts.reweight(self.len, f),
622 self.len,
623 )
624 .expect("reweight preserves EventBatch invariants");
625 reweighted.row_ids = self.row_ids.clone();
626 reweighted
627 }
628
629 pub fn slice(&self, start: usize, end: usize) -> Self {
635 assert!(start <= end);
636 assert!(end <= self.len);
637
638 if start == 0 && end == self.len {
639 return self.clone();
640 }
641
642 let mut sliced = Self::from_parts_with_len(
643 Arc::clone(&self.schema),
644 self.parts.slice(start, end),
645 end - start,
646 )
647 .expect("slice preserves EventBatch invariants");
648 sliced.row_ids = self.row_ids.as_ref().map(|ids| Arc::from(&ids[start..end]));
649 sliced
650 }
651
652 pub fn concat(batches: &[Self]) -> LadduDataResult<Self> {
659 if batches.is_empty() {
660 return Err(LadduDataError::InvalidArgument(
661 "cannot concatenate zero batches",
662 ));
663 }
664
665 let schema = Arc::clone(&batches[0].schema);
666
667 for batch in batches {
668 if schema != batch.schema {
669 return Err(LadduDataError::Schema(
670 "cannot concatenate batches with different schemas".into(),
671 ));
672 }
673 }
674
675 let parts = batches
676 .iter()
677 .map(|batch| (&batch.parts, batch.len))
678 .collect::<Vec<_>>();
679 let len = batches.iter().map(|batch| batch.len).sum();
680 let mut combined = Self::from_parts_with_len(schema, BatchParts::concat(&parts)?, len)?;
681 if batches.iter().all(|b| b.row_ids.is_some()) {
682 combined.row_ids = Some(
683 batches
684 .iter()
685 .flat_map(|b| b.row_ids.iter().flat_map(|ids| ids.iter().copied()))
686 .collect(),
687 );
688 }
689 Ok(combined)
690 }
691}
692
693fn infer_len(
694 vec4s: &[Arc<[RealVec4]>],
695 scalars: &[Arc<[f64]>],
696 columns: &[Column],
697 weight: Option<&[f64]>,
698) -> LadduDataResult<usize> {
699 let len = vec4s
700 .first()
701 .map(|c| c.len())
702 .or_else(|| scalars.first().map(|c| c.len()))
703 .or_else(|| columns.first().map(Column::len))
704 .or_else(|| weight.map(|w| w.len()))
705 .unwrap_or(0);
706
707 for col in vec4s {
708 if col.len() != len {
709 return Err(LadduDataError::Schema(
710 "inconsistent vec4 column length".into(),
711 ));
712 }
713 }
714
715 for col in scalars {
716 if col.len() != len {
717 return Err(LadduDataError::Schema(
718 "inconsistent scalar column length".into(),
719 ));
720 }
721 }
722
723 if columns.iter().any(|column| column.len() != len) {
724 return Err(LadduDataError::Schema(
725 "inconsistent typed column length".into(),
726 ));
727 }
728
729 if let Some(w) = weight
730 && w.len() != len
731 {
732 return Err(LadduDataError::Schema("inconsistent weight length".into()));
733 }
734
735 Ok(len)
736}
737
738#[derive(Copy, Clone, Debug)]
740pub struct BatchEvent<'a> {
741 batch: &'a EventBatch,
742 row: usize,
743}
744
745impl<'a> BatchEvent<'a> {
746 pub fn row(&self) -> usize {
748 self.row
749 }
750
751 pub fn batch(&self) -> &'a EventBatch {
753 self.batch
754 }
755
756 pub fn p4(&self, col: usize) -> RealVec4 {
758 self.batch.p4_at(col, self.row)
759 }
760
761 pub fn scalar(&self, col: usize) -> f64 {
763 self.batch.scalar_at(col, self.row)
764 }
765
766 pub fn column_named(&self, name: &str) -> Option<ColumnValue> {
768 self.batch
769 .column_named(name)
770 .map(|column| column.at(self.row))
771 }
772
773 pub fn weight(&self) -> f64 {
775 self.batch.weights_at(self.row)
776 }
777
778 pub fn p4_named(&self, name: &str) -> Option<RealVec4> {
780 let col = self.batch.schema.p4_index(name)?;
781 Some(self.p4(col))
782 }
783
784 pub fn scalar_named(&self, name: &str) -> Option<f64> {
786 let col = self.batch.schema.scalar_index(name)?;
787 Some(self.scalar(col))
788 }
789}
790
791#[derive(Copy, Clone, Debug)]
793pub struct Event<'a> {
794 pub(super) batch: &'a EventBatch,
795 pub(super) row: usize,
796 pub(super) weight: f64,
797}
798
799impl<'a> Event<'a> {
800 pub fn row(&self) -> usize {
802 self.row
803 }
804
805 pub fn p4(&self, col: usize) -> RealVec4 {
807 self.batch.p4_at(col, self.row)
808 }
809
810 pub fn scalar(&self, col: usize) -> f64 {
812 self.batch.scalar_at(col, self.row)
813 }
814
815 pub fn column_named(&self, name: &str) -> Option<ColumnValue> {
817 self.batch
818 .column_named(name)
819 .map(|column| column.at(self.row))
820 }
821
822 pub fn weight(&self) -> f64 {
824 self.weight
825 }
826
827 pub fn p4_named(&self, name: &str) -> Option<RealVec4> {
829 let col = self.batch.schema.p4_index(name)?;
830 Some(self.p4(col))
831 }
832
833 pub fn scalar_named(&self, name: &str) -> Option<f64> {
835 let col = self.batch.schema.scalar_index(name)?;
836 Some(self.scalar(col))
837 }
838}
839
840#[derive(Clone, Debug)]
842pub struct OwnedEvent {
843 pub p4s: Vec<RealVec4>,
845 pub scalars: Vec<f64>,
847 pub weight: Option<f64>,
849}
850
851impl OwnedEvent {
852 pub fn new(p4s: Vec<RealVec4>, scalars: Vec<f64>) -> Self {
854 Self {
855 p4s,
856 scalars,
857 weight: None,
858 }
859 }
860
861 pub fn weighted(p4s: Vec<RealVec4>, scalars: Vec<f64>, weight: f64) -> Self {
863 Self {
864 p4s,
865 scalars,
866 weight: Some(weight),
867 }
868 }
869}
870
871pub(crate) struct BatchAssembler {
873 schema: Arc<Schema>,
874 p4s: Vec<Vec<RealVec4>>,
875 scalars: Vec<Vec<f64>>,
876 columns: Vec<ColumnBuffer>,
877 weights: WeightAssembler,
878 len: usize,
879}
880
881impl BatchAssembler {
882 pub(crate) fn new(schema: Arc<Schema>, capacity: usize) -> Self {
883 let p4s = (0..schema.n_p4s())
884 .map(|_| Vec::with_capacity(capacity))
885 .collect();
886 let scalars = (0..schema.n_scalars())
887 .map(|_| Vec::with_capacity(capacity))
888 .collect();
889
890 Self {
891 columns: schema
892 .columns()
893 .iter()
894 .map(|(_, dtype)| ColumnBuffer::new(*dtype, capacity))
895 .collect(),
896 schema,
897 p4s,
898 scalars,
899 weights: WeightAssembler::default(),
900 len: 0,
901 }
902 }
903
904 pub(crate) fn with_weight_mode(
905 schema: Arc<Schema>,
906 capacity: usize,
907 explicit_weights: bool,
908 ) -> Self {
909 let mut assembler = Self::new(schema, capacity);
910 if explicit_weights {
911 assembler.weights = WeightAssembler::Explicit(Vec::with_capacity(capacity));
912 }
913 assembler
914 }
915
916 pub(crate) fn from_columns(
917 schema: Arc<Schema>,
918 p4s: Vec<Arc<[RealVec4]>>,
919 scalars: Vec<Arc<[f64]>>,
920 weights: Option<Arc<[f64]>>,
921 ) -> LadduDataResult<EventBatch> {
922 EventBatch::from_parts(
923 schema,
924 BatchParts::from_columns(p4s, scalars, Weights::from_option(weights)),
925 )
926 }
927
928 fn push_owned(&mut self, event: OwnedEvent) -> LadduDataResult<()> {
929 if self.schema.n_columns() != 0 {
930 return Err(LadduDataError::Unsupported(
931 "owned float events cannot populate typed columns; construct typed columnar batches",
932 ));
933 }
934 if event.p4s.len() != self.schema.n_p4s() {
935 return Err(LadduDataError::Schema(
936 "wrong number of event vec4 values".into(),
937 ));
938 }
939
940 if event.scalars.len() != self.schema.n_scalars() {
941 return Err(LadduDataError::Schema(
942 "wrong number of event scalar values".into(),
943 ));
944 }
945
946 self.weights.push(event.weight, self.len)?;
947
948 for (col, value) in event.p4s.into_iter().enumerate() {
949 self.p4s[col].push(value);
950 }
951 for (col, value) in event.scalars.into_iter().enumerate() {
952 self.scalars[col].push(value);
953 }
954
955 self.len += 1;
956 Ok(())
957 }
958
959 pub(crate) fn push_borrowed(
960 &mut self,
961 event: Event<'_>,
962 explicit_weight: bool,
963 ) -> LadduDataResult<()> {
964 if self.schema.columns() != event.batch.schema().columns() {
965 return Err(LadduDataError::Schema(
966 "typed event columns do not match schema".into(),
967 ));
968 }
969 if event.batch.schema().n_p4s() != self.schema.n_p4s() {
970 return Err(LadduDataError::Schema(
971 "wrong number of event vec4 values".into(),
972 ));
973 }
974
975 if event.batch.schema().n_scalars() != self.schema.n_scalars() {
976 return Err(LadduDataError::Schema(
977 "wrong number of event scalar values".into(),
978 ));
979 }
980
981 self.weights
982 .push(explicit_weight.then_some(event.weight()), self.len)?;
983
984 for col in 0..self.schema.n_p4s() {
985 self.p4s[col].push(event.p4(col));
986 }
987 for col in 0..self.schema.n_scalars() {
988 self.scalars[col].push(event.scalar(col));
989 }
990
991 for (index, column) in self.columns.iter_mut().enumerate() {
992 column.push(event.batch.column(index).at(event.row))?;
993 }
994
995 self.len += 1;
996 Ok(())
997 }
998
999 pub(crate) fn finish(self) -> LadduDataResult<EventBatch> {
1000 let mut parts = BatchParts::from_columns(
1001 self.p4s.into_iter().map(Arc::from).collect(),
1002 self.scalars.into_iter().map(Arc::from).collect(),
1003 self.weights.finish(),
1004 );
1005 parts.columns = self.columns.into_iter().map(ColumnBuffer::finish).collect();
1006 EventBatch::from_parts_with_len(self.schema, parts, self.len)
1007 }
1008}
1009
1010pub struct EventBatchBuilder {
1012 assembler: BatchAssembler,
1013}
1014
1015impl EventBatchBuilder {
1016 pub fn new(schema: Arc<Schema>) -> Self {
1018 Self::with_capacity(schema, 0)
1019 }
1020
1021 pub fn with_capacity(schema: Arc<Schema>, capacity: usize) -> Self {
1023 Self {
1024 assembler: BatchAssembler::new(schema, capacity),
1025 }
1026 }
1027
1028 pub fn push<P, S>(&mut self, p4s: P, scalars: S) -> LadduDataResult<&mut Self>
1035 where
1036 P: IntoIterator<Item = RealVec4>,
1037 S: IntoIterator<Item = f64>,
1038 {
1039 self.push_event(OwnedEvent::new(
1040 p4s.into_iter().collect(),
1041 scalars.into_iter().collect(),
1042 ))
1043 }
1044
1045 pub fn push_weighted<P, S>(
1052 &mut self,
1053 p4s: P,
1054 scalars: S,
1055 weight: f64,
1056 ) -> LadduDataResult<&mut Self>
1057 where
1058 P: IntoIterator<Item = RealVec4>,
1059 S: IntoIterator<Item = f64>,
1060 {
1061 self.push_event(OwnedEvent::weighted(
1062 p4s.into_iter().collect(),
1063 scalars.into_iter().collect(),
1064 weight,
1065 ))
1066 }
1067
1068 pub fn push_event(&mut self, event: OwnedEvent) -> LadduDataResult<&mut Self> {
1075 self.assembler.push_owned(event)?;
1076 Ok(self)
1077 }
1078
1079 pub fn extend<I>(&mut self, events: I) -> LadduDataResult<&mut Self>
1086 where
1087 I: IntoIterator<Item = OwnedEvent>,
1088 {
1089 for event in events {
1090 self.push_event(event)?;
1091 }
1092
1093 Ok(self)
1094 }
1095
1096 pub fn finish(self) -> LadduDataResult<EventBatch> {
1103 self.assembler.finish()
1104 }
1105}
1106
1107#[cfg(test)]
1108mod tests {
1109 use super::*;
1110
1111 #[test]
1112 fn column_free_batches_preserve_counts_and_global_identities() -> LadduDataResult<()> {
1113 let schema = Arc::new(Schema::new(
1114 Vec::<String>::new(),
1115 Vec::<String>::new(),
1116 false,
1117 )?);
1118 let batch = EventBatch::new_with_len(schema, vec![], vec![], None, 3)?
1119 .with_row_ids(Arc::from([2, 5, 8]))?;
1120 assert_eq!(batch.resident_bytes(), 3 * std::mem::size_of::<u64>());
1121 assert_eq!(batch.reweight(|_, w| w).len(), 3);
1122 assert_eq!(batch.select(&[2, 0]).row_ids(), Some([8, 2].as_slice()));
1123 let rebuilt = EventBatch::concat(&[batch.slice(0, 1), batch.slice(1, 3)])?;
1124 assert_eq!(rebuilt.row_ids(), Some([2, 5, 8].as_slice()));
1125 assert!(batch.clone().with_row_ids(Arc::from([1])).is_err());
1126 Ok(())
1127 }
1128
1129 fn v(x: f64) -> RealVec4 {
1130 RealVec4 {
1131 e: x + 0.3,
1132 px: x,
1133 py: x + 0.1,
1134 pz: x + 0.2,
1135 }
1136 }
1137
1138 fn schema_with_weight() -> Arc<Schema> {
1139 Arc::new(Schema::new(["p"], ["x"], true).unwrap())
1140 }
1141
1142 fn weighted_batch(start: usize, len: usize) -> EventBatch {
1143 let schema = schema_with_weight();
1144
1145 let events = (start..start + len)
1146 .map(|i| OwnedEvent::weighted(vec![v(i as f64)], vec![i as f64], 10.0 + i as f64));
1147
1148 EventBatch::from_events(schema, events).unwrap()
1149 }
1150
1151 fn scalar_values(batch: &EventBatch) -> Vec<f64> {
1152 batch.scalar_column(0).to_vec()
1153 }
1154
1155 #[test]
1156 fn bound_columns_reuse_schema_resolution_for_row_access() {
1157 let batch = weighted_batch(3, 2);
1158 let scalar = batch.schema().bind_scalar("x").unwrap();
1159 let p4 = batch.schema().bind_p4("p").unwrap();
1160
1161 assert_eq!(batch.scalar_column_bound(&scalar).unwrap(), &[3.0, 4.0]);
1162 assert_eq!(batch.p4_column_bound(&p4).unwrap(), &[v(3.0), v(4.0)]);
1163
1164 let other_schema = Arc::new(Schema::new(["other"], ["x"], true).unwrap());
1165 let other = other_schema.bind_scalar("x").unwrap();
1166 assert!(matches!(
1167 batch.scalar_column_bound(&other),
1168 Err(LadduDataError::Schema(message)) if message == "column binding belongs to a different schema"
1169 ));
1170 }
1171
1172 #[test]
1173 fn event_batch_rejects_shape_mismatches_and_builder_rejects_mixed_weights() {
1174 let schema = schema_with_weight();
1175
1176 let bad_vec4_count = EventBatch::new(
1177 Arc::clone(&schema),
1178 vec![],
1179 vec![Arc::from([1.0, 2.0])],
1180 Some(Arc::from([1.0, 2.0])),
1181 );
1182
1183 assert!(matches!(bad_vec4_count, Err(LadduDataError::Schema(_))));
1184
1185 let bad_lengths = EventBatch::new(
1186 Arc::clone(&schema),
1187 vec![Arc::from([v(1.0), v(2.0)])],
1188 vec![Arc::from([1.0])],
1189 Some(Arc::from([1.0, 2.0])),
1190 );
1191
1192 assert!(matches!(bad_lengths, Err(LadduDataError::Schema(_))));
1193
1194 let mut builder = EventBatchBuilder::new(schema);
1195 builder.push([v(1.0)], [1.0]).unwrap();
1196
1197 let mixed = builder.push_weighted([v(2.0)], [2.0], 2.0);
1198 assert!(matches!(mixed, Err(LadduDataError::InvalidArgument(_))));
1199 }
1200
1201 #[test]
1202 fn select_slice_filter_reweight_and_concat_preserve_columns_and_weight_semantics() {
1203 let weighted = weighted_batch(0, 4);
1204 let selected = weighted.select(&[3, 1]);
1205
1206 assert_eq!(scalar_values(&selected), vec![3.0, 1.0]);
1207 assert_eq!(selected.weights_column().unwrap(), &[13.0, 11.0]);
1208 assert_eq!(selected.p4_at(0, 0).px, 3.0);
1209 assert_eq!(selected.p4_at(0, 1).e, 1.3);
1210
1211 let sliced = weighted.slice(1, 3);
1212 assert_eq!(scalar_values(&sliced), vec![1.0, 2.0]);
1213 assert_eq!(sliced.weights_column().unwrap(), &[11.0, 12.0]);
1214
1215 let filtered = weighted.filter(|ev| ev.scalar(0) >= 2.0);
1216 assert_eq!(scalar_values(&filtered), vec![2.0, 3.0]);
1217
1218 let reweighted = filtered.reweight(|i, w| w + 100.0 + i as f64);
1219 assert_eq!(reweighted.weights_column().unwrap(), &[112.0, 114.0]);
1220
1221 let schema = schema_with_weight();
1222
1223 let unweighted_with_weight_schema = EventBatch::from_events(
1224 Arc::clone(&schema),
1225 [
1226 OwnedEvent::new(vec![v(100.0)], vec![100.0]),
1227 OwnedEvent::new(vec![v(101.0)], vec![101.0]),
1228 ],
1229 )
1230 .unwrap();
1231
1232 let weighted_tail = EventBatch::from_events(
1233 schema,
1234 [
1235 OwnedEvent::weighted(vec![v(200.0)], vec![200.0], 5.0),
1236 OwnedEvent::weighted(vec![v(201.0)], vec![201.0], 6.0),
1237 ],
1238 )
1239 .unwrap();
1240
1241 let concatenated =
1242 EventBatch::concat(&[unweighted_with_weight_schema, weighted_tail]).unwrap();
1243
1244 assert_eq!(
1245 scalar_values(&concatenated),
1246 vec![100.0, 101.0, 200.0, 201.0]
1247 );
1248 assert_eq!(
1249 concatenated.weights_column().unwrap(),
1250 &[1.0, 1.0, 5.0, 6.0]
1251 );
1252 }
1253
1254 #[test]
1255 fn implicit_unit_weights_survive_assembly_and_row_transforms() {
1256 let schema = Arc::new(Schema::new(["p"], ["x"], false).unwrap());
1257 let batch = EventBatch::from_events(
1258 Arc::clone(&schema),
1259 (0..3).map(|i| OwnedEvent::new(vec![v(i as f64)], vec![i as f64])),
1260 )
1261 .unwrap();
1262
1263 assert!(batch.weights_column().is_none());
1264 assert_eq!(batch.weights_at(2), 1.0);
1265
1266 let selected = batch.select(&[2, 0]);
1267 let sliced = batch.slice(1, 3);
1268 let filtered = batch.filter(|event| event.scalar(0) > 0.0);
1269 let concatenated = EventBatch::concat(&[selected, sliced]).unwrap();
1270
1271 assert!(filtered.weights_column().is_none());
1272 assert!(concatenated.weights_column().is_none());
1273 assert_eq!(concatenated.weights_at(3), 1.0);
1274
1275 let reweighted = batch.reweight(|row, weight| weight + row as f64);
1276 assert_eq!(reweighted.weights_column().unwrap(), &[1.0, 2.0, 3.0]);
1277 }
1278
1279 #[test]
1280 fn shared_assembler_preserves_weight_mode_and_rejects_transitions() {
1281 let schema = Arc::new(Schema::new(["p"], ["x"], true).unwrap());
1282 let source = EventBatch::from_events(
1283 Arc::clone(&schema),
1284 [
1285 OwnedEvent::weighted(vec![v(1.0)], vec![1.0], 2.0),
1286 OwnedEvent::weighted(vec![v(2.0)], vec![2.0], 3.0),
1287 ],
1288 )
1289 .unwrap();
1290
1291 let mut explicit = BatchAssembler::new(Arc::clone(&schema), 2);
1292 let first = Event {
1293 batch: &source,
1294 row: 0,
1295 weight: source.weights_at(0),
1296 };
1297 explicit.push_borrowed(first, true).unwrap();
1298 let second = Event {
1299 batch: &source,
1300 row: 1,
1301 weight: source.weights_at(1),
1302 };
1303 let transition = explicit.push_borrowed(second, false);
1304 assert!(matches!(
1305 transition,
1306 Err(LadduDataError::InvalidArgument(_))
1307 ));
1308 let explicit = explicit.finish().unwrap();
1309 assert_eq!(explicit.len(), 1);
1310 assert_eq!(explicit.weights_column().unwrap(), &[2.0]);
1311
1312 let unweighted = EventBatch::from_events(
1313 Arc::clone(&schema),
1314 [OwnedEvent::new(vec![v(3.0)], vec![3.0])],
1315 )
1316 .unwrap();
1317 let mut implicit = BatchAssembler::new(schema, 1);
1318 let event = Event {
1319 batch: &unweighted,
1320 row: 0,
1321 weight: unweighted.weights_at(0),
1322 };
1323 implicit.push_borrowed(event, false).unwrap();
1324 let implicit = implicit.finish().unwrap();
1325 assert!(implicit.weights_column().is_none());
1326 assert_eq!(implicit.weights_at(0), 1.0);
1327 }
1328
1329 #[test]
1330 fn assembly_table_covers_column_shapes_and_observable_sharing() {
1331 let cases = [
1332 (
1333 Arc::new(Schema::new(Vec::<&str>::new(), Vec::<&str>::new(), false).unwrap()),
1334 false,
1335 ),
1336 (
1337 Arc::new(Schema::new(["p"], Vec::<&str>::new(), false).unwrap()),
1338 false,
1339 ),
1340 (
1341 Arc::new(Schema::new(Vec::<&str>::new(), ["x"], false).unwrap()),
1342 false,
1343 ),
1344 (
1345 Arc::new(Schema::new(Vec::<&str>::new(), Vec::<&str>::new(), true).unwrap()),
1346 true,
1347 ),
1348 (Arc::new(Schema::new(["p"], ["x"], true).unwrap()), true),
1349 ];
1350
1351 for (schema, weighted) in cases {
1352 let events = (0..2).map(|i| {
1353 let p4s = if schema.n_p4s() == 0 {
1354 Vec::new()
1355 } else {
1356 vec![v(i as f64)]
1357 };
1358 let scalars = if schema.n_scalars() == 0 {
1359 Vec::new()
1360 } else {
1361 vec![i as f64]
1362 };
1363 if weighted {
1364 OwnedEvent::weighted(p4s, scalars, 2.0 + i as f64)
1365 } else {
1366 OwnedEvent::new(p4s, scalars)
1367 }
1368 });
1369 let batch = EventBatch::from_events(Arc::clone(&schema), events).unwrap();
1370 assert_eq!(batch.len(), 2);
1371 assert_eq!(batch.weights_column().is_some(), weighted);
1372
1373 let selected = batch.select(&(0..batch.len()).collect::<Vec<_>>());
1374 assert_eq!(selected.len(), batch.len());
1375 for col in 0..schema.n_p4s() {
1376 assert_eq!(selected.vec4_column(col), batch.vec4_column(col));
1377 }
1378 for col in 0..schema.n_scalars() {
1379 assert_eq!(selected.scalar_column(col), batch.scalar_column(col));
1380 }
1381 assert_eq!(selected.weights_column(), batch.weights_column());
1382
1383 let concatenated = EventBatch::concat(&[batch.slice(0, 1), batch.slice(1, 2)]).unwrap();
1384 assert_eq!(concatenated.len(), batch.len());
1385 for col in 0..schema.n_p4s() {
1386 assert_eq!(concatenated.vec4_column(col), batch.vec4_column(col));
1387 }
1388 for col in 0..schema.n_scalars() {
1389 assert_eq!(concatenated.scalar_column(col), batch.scalar_column(col));
1390 }
1391 assert_eq!(concatenated.weights_column(), batch.weights_column());
1392
1393 if schema.n_p4s() > 0 || schema.n_scalars() > 0 {
1394 let reweighted = batch.reweight(|_, weight| weight + 1.0);
1395 if schema.n_p4s() > 0 {
1396 assert_eq!(
1397 reweighted.vec4_column(0).as_ptr(),
1398 batch.vec4_column(0).as_ptr()
1399 );
1400 }
1401 if schema.n_scalars() > 0 {
1402 assert_eq!(
1403 reweighted.scalar_column(0).as_ptr(),
1404 batch.scalar_column(0).as_ptr()
1405 );
1406 }
1407 }
1408 }
1409 }
1410}