1use serde::{Deserialize, Serialize};
2
3use crate::{
4 LadduPhysicsError, LadduPhysicsResult,
5 binning::{BinningAxis, FinalUpperEdge, bin_shape, flat_bin_index},
6 histogram::{HistogramUncertaintyStatus, uncertainty_is_available},
7};
8
9#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
11pub struct JointHistogramDiagnostics {
12 nonfinite_count: u64,
13 nonfinite_weight: f64,
14 out_of_range_count: u64,
15 out_of_range_weight: f64,
16}
17
18impl JointHistogramDiagnostics {
19 pub fn nonfinite_count(&self) -> u64 {
21 self.nonfinite_count
22 }
23 pub fn nonfinite_weight(&self) -> f64 {
25 self.nonfinite_weight
26 }
27 pub fn out_of_range_count(&self) -> u64 {
29 self.out_of_range_count
30 }
31 pub fn out_of_range_weight(&self) -> f64 {
33 self.out_of_range_weight
34 }
35}
36
37#[derive(Clone, Debug, PartialEq, Serialize)]
42pub struct JointHistogram {
43 axes: Vec<Vec<f64>>,
44 #[serde(skip)]
45 bin_axes: Vec<BinningAxis>,
46 shape: Vec<usize>,
47 values: Vec<f64>,
48 sum_squared_weights: Vec<f64>,
49 diagnostics: JointHistogramDiagnostics,
50 #[serde(default, skip_serializing_if = "uncertainty_is_available")]
51 uncertainty_status: HistogramUncertaintyStatus,
52 #[serde(skip)]
53 value_corrections: Vec<f64>,
54 #[serde(skip)]
55 squared_weight_corrections: Vec<f64>,
56}
57
58#[derive(Deserialize)]
59struct SerializedJointHistogram {
60 axes: Vec<Vec<f64>>,
61 shape: Vec<usize>,
62 values: Vec<f64>,
63 sum_squared_weights: Vec<f64>,
64 diagnostics: JointHistogramDiagnostics,
65 #[serde(default)]
66 uncertainty_status: HistogramUncertaintyStatus,
67}
68
69impl<'de> Deserialize<'de> for JointHistogram {
70 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
71 where
72 D: serde::Deserializer<'de>,
73 {
74 let serialized = SerializedJointHistogram::deserialize(deserializer)?;
75 let mut histogram = Self::empty(serialized.axes).map_err(serde::de::Error::custom)?;
76 if histogram.shape != serialized.shape
77 || histogram.values.len() != serialized.values.len()
78 || histogram.sum_squared_weights.len() != serialized.sum_squared_weights.len()
79 {
80 return Err(serde::de::Error::custom(
81 "joint histogram shape and flattened accumulator lengths are inconsistent",
82 ));
83 }
84 if serialized
85 .values
86 .iter()
87 .chain(&serialized.sum_squared_weights)
88 .any(|value| !value.is_finite())
89 || serialized
90 .sum_squared_weights
91 .iter()
92 .any(|value| *value < 0.0)
93 || !serialized.diagnostics.nonfinite_weight.is_finite()
94 || !serialized.diagnostics.out_of_range_weight.is_finite()
95 {
96 return Err(serde::de::Error::custom(
97 "joint histogram accumulators and diagnostic weights must be finite",
98 ));
99 }
100 histogram.values = serialized.values;
101 histogram.sum_squared_weights = serialized.sum_squared_weights;
102 histogram.diagnostics = serialized.diagnostics;
103 histogram.uncertainty_status = serialized.uncertainty_status;
104 Ok(histogram)
105 }
106}
107
108impl JointHistogram {
109 pub fn empty(axes: Vec<Vec<f64>>) -> LadduPhysicsResult<Self> {
116 if axes.is_empty() {
117 return Err(LadduPhysicsError::invalid_length(
118 "joint histogram axes",
119 "at least 1",
120 0,
121 ));
122 }
123 let bin_axes = axes
124 .iter()
125 .enumerate()
126 .map(|(axis, edges)| {
127 BinningAxis::new(edges.iter().copied()).map_err(|_| {
128 LadduPhysicsError::invalid_relation(format!(
129 "joint histogram axis {axis} edges must contain at least two finite, strictly increasing values"
130 ))
131 })
132 })
133 .collect::<LadduPhysicsResult<Vec<_>>>()?;
134 let shape = bin_shape(&bin_axes);
135 let bins = checked_bin_count(&shape)?;
136 let values = zeroed(bins)?;
137 let sum_squared_weights = zeroed(bins)?;
138 let value_corrections = zeroed(bins)?;
139 let squared_weight_corrections = zeroed(bins)?;
140 Ok(Self {
141 axes,
142 bin_axes,
143 shape,
144 values,
145 sum_squared_weights,
146 diagnostics: JointHistogramDiagnostics::default(),
147 uncertainty_status: HistogramUncertaintyStatus::Available,
148 value_corrections,
149 squared_weight_corrections,
150 })
151 }
152
153 pub fn axes(&self) -> &[Vec<f64>] {
155 &self.axes
156 }
157 pub fn shape(&self) -> &[usize] {
159 &self.shape
160 }
161 pub fn values(&self) -> &[f64] {
163 &self.values
164 }
165 pub fn sum_squared_weights(&self) -> &[f64] {
167 &self.sum_squared_weights
168 }
169 pub fn errors(&self) -> Vec<f64> {
171 self.sum_squared_weights
172 .iter()
173 .map(|value| value.sqrt())
174 .collect()
175 }
176 pub fn reported_errors(&self) -> Option<Vec<f64>> {
178 (self.uncertainty_status == HistogramUncertaintyStatus::Available).then(|| self.errors())
179 }
180 pub fn uncertainty_status(&self) -> HistogramUncertaintyStatus {
182 self.uncertainty_status
183 }
184 pub fn diagnostics(&self) -> &JointHistogramDiagnostics {
186 &self.diagnostics
187 }
188
189 pub fn fill_weighted(&mut self, coordinates: &[f64], weight: f64) -> LadduPhysicsResult<()> {
196 if coordinates.len() != self.axes.len() {
197 return Err(LadduPhysicsError::invalid_length(
198 "joint histogram coordinates",
199 self.axes.len().to_string(),
200 coordinates.len(),
201 ));
202 }
203 if coordinates.iter().any(|value| !value.is_finite()) || !weight.is_finite() {
204 self.diagnostics.nonfinite_count += 1;
205 if weight.is_finite() {
206 let next_weight = self.diagnostics.nonfinite_weight + weight;
207 if !next_weight.is_finite() {
208 return Err(LadduPhysicsError::invalid_relation(
209 "joint histogram nonfinite diagnostic weight overflow",
210 ));
211 }
212 self.diagnostics.nonfinite_weight = next_weight;
213 }
214 return Ok(());
215 }
216 let squared_weight = weight * weight;
217 if !squared_weight.is_finite() {
218 return Err(LadduPhysicsError::invalid_value(
219 "joint histogram squared weight",
220 "finite",
221 squared_weight,
222 ));
223 }
224 let Some(flat) = flat_bin_index(&self.bin_axes, coordinates, FinalUpperEdge::Exclusive)
225 else {
226 self.diagnostics.out_of_range_count += 1;
227 let next_weight = self.diagnostics.out_of_range_weight + weight;
228 if !next_weight.is_finite() {
229 return Err(LadduPhysicsError::invalid_relation(
230 "joint histogram out-of-range weight overflow",
231 ));
232 }
233 self.diagnostics.out_of_range_weight = next_weight;
234 return Ok(());
235 };
236 if self.value_corrections.len() != self.values.len() {
237 self.value_corrections.resize(self.values.len(), 0.0);
238 }
239 if self.squared_weight_corrections.len() != self.values.len() {
240 self.squared_weight_corrections
241 .resize(self.values.len(), 0.0);
242 }
243 compensated_add(
244 &mut self.values[flat],
245 &mut self.value_corrections[flat],
246 weight,
247 );
248 compensated_add(
249 &mut self.sum_squared_weights[flat],
250 &mut self.squared_weight_corrections[flat],
251 squared_weight,
252 );
253 if !self.values[flat].is_finite() || !self.sum_squared_weights[flat].is_finite() {
254 return Err(LadduPhysicsError::invalid_relation(
255 "joint histogram fill produced a non-finite accumulator",
256 ));
257 }
258 Ok(())
259 }
260
261 pub fn merge(&mut self, other: &Self) -> LadduPhysicsResult<()> {
268 if self.axes != other.axes || self.shape != other.shape {
269 return Err(LadduPhysicsError::invalid_relation(
270 "joint histogram merge requires identical ordered axes and shape",
271 ));
272 }
273 let mut merged = self.clone();
274 if other.uncertainty_status == HistogramUncertaintyStatus::UnavailableCovariance {
275 merged.uncertainty_status = HistogramUncertaintyStatus::UnavailableCovariance;
276 }
277 merged.value_corrections.resize(merged.values.len(), 0.0);
278 merged
279 .squared_weight_corrections
280 .resize(merged.values.len(), 0.0);
281 for index in 0..merged.values.len() {
282 compensated_add(
283 &mut merged.values[index],
284 &mut merged.value_corrections[index],
285 other.values[index],
286 );
287 compensated_add(
288 &mut merged.sum_squared_weights[index],
289 &mut merged.squared_weight_corrections[index],
290 other.sum_squared_weights[index],
291 );
292 }
293 merged.diagnostics.nonfinite_count = merged
294 .diagnostics
295 .nonfinite_count
296 .checked_add(other.diagnostics.nonfinite_count)
297 .ok_or_else(|| {
298 LadduPhysicsError::invalid_relation("joint histogram diagnostic count overflow")
299 })?;
300 merged.diagnostics.out_of_range_count = merged
301 .diagnostics
302 .out_of_range_count
303 .checked_add(other.diagnostics.out_of_range_count)
304 .ok_or_else(|| {
305 LadduPhysicsError::invalid_relation("joint histogram diagnostic count overflow")
306 })?;
307 merged.diagnostics.nonfinite_weight += other.diagnostics.nonfinite_weight;
308 merged.diagnostics.out_of_range_weight += other.diagnostics.out_of_range_weight;
309 if merged
310 .values
311 .iter()
312 .chain(&merged.sum_squared_weights)
313 .any(|value| !value.is_finite())
314 || !merged.diagnostics.nonfinite_weight.is_finite()
315 || !merged.diagnostics.out_of_range_weight.is_finite()
316 {
317 return Err(LadduPhysicsError::invalid_relation(
318 "joint histogram merge produced a non-finite accumulator",
319 ));
320 }
321 *self = merged;
322 Ok(())
323 }
324
325 pub fn add(&self, other: &Self) -> LadduPhysicsResult<Self> {
330 self.combine(other, 1.0, false)
331 }
332 pub fn add_independent(&self, other: &Self) -> LadduPhysicsResult<Self> {
337 self.combine(other, 1.0, true)
338 }
339 pub fn subtract(&self, other: &Self) -> LadduPhysicsResult<Self> {
344 self.combine(other, -1.0, false)
345 }
346 pub fn subtract_independent(&self, other: &Self) -> LadduPhysicsResult<Self> {
351 self.combine(other, -1.0, true)
352 }
353
354 fn combine(
355 &self,
356 other: &Self,
357 sign: f64,
358 assert_independent: bool,
359 ) -> LadduPhysicsResult<Self> {
360 if self.axes != other.axes || self.shape != other.shape {
361 return Err(LadduPhysicsError::invalid_relation(
362 "joint histogram arithmetic requires identical ordered axes and shape",
363 ));
364 }
365 let mut result = self.clone();
366 result.value_corrections.resize(result.values.len(), 0.0);
367 result
368 .squared_weight_corrections
369 .resize(result.values.len(), 0.0);
370 for index in 0..result.values.len() {
371 compensated_add(
372 &mut result.values[index],
373 &mut result.value_corrections[index],
374 sign * other.values[index],
375 );
376 compensated_add(
377 &mut result.sum_squared_weights[index],
378 &mut result.squared_weight_corrections[index],
379 other.sum_squared_weights[index],
380 );
381 }
382 result.diagnostics.nonfinite_count = result
383 .diagnostics
384 .nonfinite_count
385 .checked_add(other.diagnostics.nonfinite_count)
386 .ok_or_else(|| {
387 LadduPhysicsError::invalid_relation("joint histogram diagnostic count overflow")
388 })?;
389 result.diagnostics.out_of_range_count = result
390 .diagnostics
391 .out_of_range_count
392 .checked_add(other.diagnostics.out_of_range_count)
393 .ok_or_else(|| {
394 LadduPhysicsError::invalid_relation("joint histogram diagnostic count overflow")
395 })?;
396 result.diagnostics.nonfinite_weight += sign * other.diagnostics.nonfinite_weight;
397 result.diagnostics.out_of_range_weight += sign * other.diagnostics.out_of_range_weight;
398 result.uncertainty_status = if self.uncertainty_status
399 == HistogramUncertaintyStatus::Available
400 && other.uncertainty_status == HistogramUncertaintyStatus::Available
401 && assert_independent
402 {
403 HistogramUncertaintyStatus::Available
404 } else {
405 HistogramUncertaintyStatus::UnavailableCovariance
406 };
407 if result
408 .values
409 .iter()
410 .chain(&result.sum_squared_weights)
411 .any(|value| !value.is_finite())
412 || !result.diagnostics.nonfinite_weight.is_finite()
413 || !result.diagnostics.out_of_range_weight.is_finite()
414 {
415 return Err(LadduPhysicsError::invalid_relation(
416 "joint histogram arithmetic produced a non-finite accumulator",
417 ));
418 }
419 Ok(result)
420 }
421
422 pub fn scaled(&self, factor: f64) -> LadduPhysicsResult<Self> {
427 if !factor.is_finite() {
428 return Err(LadduPhysicsError::invalid_value(
429 "joint histogram scale",
430 "finite",
431 factor,
432 ));
433 }
434 let mut result = self.clone();
435 let variance_scale = factor * factor;
436 for value in &mut result.values {
437 *value *= factor;
438 }
439 for value in &mut result.sum_squared_weights {
440 *value *= variance_scale;
441 }
442 result.diagnostics.nonfinite_weight *= factor;
443 result.diagnostics.out_of_range_weight *= factor;
444 result.value_corrections.fill(0.0);
445 result.squared_weight_corrections.fill(0.0);
446 if result
447 .values
448 .iter()
449 .chain(&result.sum_squared_weights)
450 .any(|value| !value.is_finite())
451 || !result.diagnostics.nonfinite_weight.is_finite()
452 || !result.diagnostics.out_of_range_weight.is_finite()
453 {
454 return Err(LadduPhysicsError::invalid_relation(
455 "joint histogram scaling produced a non-finite accumulator",
456 ));
457 }
458 Ok(result)
459 }
460}
461
462fn zeroed(len: usize) -> LadduPhysicsResult<Vec<f64>> {
463 let mut values = Vec::new();
464 values.try_reserve_exact(len).map_err(|_| {
465 LadduPhysicsError::invalid_relation(format!(
466 "joint histogram shape with {len} bins cannot be allocated"
467 ))
468 })?;
469 values.resize(len, 0.0);
470 Ok(values)
471}
472
473fn checked_bin_count(shape: &[usize]) -> LadduPhysicsResult<usize> {
474 shape.iter().try_fold(1usize, |bins, axis_bins| {
475 bins.checked_mul(*axis_bins).ok_or_else(|| {
476 LadduPhysicsError::invalid_relation("joint histogram shape overflows usize")
477 })
478 })
479}
480
481fn compensated_add(sum: &mut f64, correction: &mut f64, value: f64) {
482 let adjusted = value - *correction;
483 let next = *sum + adjusted;
484 *correction = (next - *sum) - adjusted;
485 *sum = next;
486}
487
488#[cfg(test)]
489mod tests {
490 use super::*;
491
492 #[test]
493 fn arithmetic_distinguishes_proven_independence_from_shared_sources() {
494 let mut left = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
495 left.fill_weighted(&[0.5], 2.0).unwrap();
496 let mut independent = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
497 independent.fill_weighted(&[0.5], 3.0).unwrap();
498
499 let conservative_sum = left.add(&independent).unwrap();
500 assert_eq!(conservative_sum.values(), [5.0]);
501 assert_eq!(
502 conservative_sum.uncertainty_status(),
503 HistogramUncertaintyStatus::UnavailableCovariance
504 );
505 assert!(conservative_sum.reported_errors().is_none());
506
507 let asserted = left.add_independent(&independent).unwrap();
508 assert_eq!(asserted.reported_errors().unwrap(), [13.0_f64.sqrt()]);
509
510 let correlated_sum = left.add(&left).unwrap();
511 assert_eq!(correlated_sum.values(), [4.0]);
512 assert_eq!(
513 correlated_sum.uncertainty_status(),
514 HistogramUncertaintyStatus::UnavailableCovariance
515 );
516 assert!(correlated_sum.reported_errors().is_none());
517 }
518
519 #[test]
520 fn arithmetic_scales_serializes_and_rejects_incompatible_axes() {
521 let mut histogram = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
522 histogram.fill_weighted(&[0.5], -2.0).unwrap();
523 let scaled = histogram.scaled(-3.0).unwrap();
524 assert_eq!(scaled.values(), [6.0]);
525 assert_eq!(scaled.sum_squared_weights(), [36.0]);
526
527 let unavailable = histogram.subtract(&histogram).unwrap();
528 let restored: JointHistogram =
529 serde_json::from_str(&serde_json::to_string(&unavailable).unwrap()).unwrap();
530 assert_eq!(
531 restored.uncertainty_status(),
532 HistogramUncertaintyStatus::UnavailableCovariance
533 );
534 assert!(restored.reported_errors().is_none());
535
536 let incompatible = JointHistogram::empty(vec![vec![0.0, 2.0]]).unwrap();
537 assert!(histogram.add(&incompatible).is_err());
538 assert!(histogram.scaled(f64::NAN).is_err());
539 }
540
541 #[test]
542 fn repeated_arithmetic_uses_deterministic_compensated_accumulation() {
543 let mut large = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
544 large.fill_weighted(&[0.5], 1.0e16).unwrap();
545 let mut unit = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
546 unit.fill_weighted(&[0.5], 1.0).unwrap();
547
548 let result = large
549 .add_independent(&unit)
550 .unwrap()
551 .add_independent(&unit)
552 .unwrap();
553
554 assert_eq!(result.values(), [1.0e16 + 2.0]);
555 }
556
557 #[test]
558 fn validates_shape_and_merges_disjoint_fills() {
559 let axes = vec![vec![0.0, 1.0, 2.0], vec![0.0, 5.0, 10.0]];
560 let mut left = JointHistogram::empty(axes.clone()).unwrap();
561 let mut right = JointHistogram::empty(axes).unwrap();
562 left.fill_weighted(&[0.5, 7.0], -2.0).unwrap();
563 right.fill_weighted(&[1.5, 2.0], 3.0).unwrap();
564 left.merge(&right).unwrap();
565 assert_eq!(left.values(), [0.0, -2.0, 3.0, 0.0]);
566 assert_eq!(left.errors(), [0.0, 2.0, 3.0, 0.0]);
567 }
568
569 #[test]
570 fn final_upper_edges_are_out_of_range_and_nonfinite_takes_precedence() {
571 let mut histogram = JointHistogram::empty(vec![vec![0.0, 1.0], vec![0.0, 1.0]]).unwrap();
572 histogram.fill_weighted(&[1.0, 0.5], -2.0).unwrap();
573 histogram.fill_weighted(&[f64::NAN, 2.0], 3.0).unwrap();
574
575 assert_eq!(histogram.values(), [0.0]);
576 assert_eq!(histogram.diagnostics().out_of_range_count(), 1);
577 assert_eq!(histogram.diagnostics().out_of_range_weight(), -2.0);
578 assert_eq!(histogram.diagnostics().nonfinite_count(), 1);
579 assert_eq!(histogram.diagnostics().nonfinite_weight(), 3.0);
580 }
581
582 #[test]
583 fn nonuniform_lower_and_internal_boundaries_follow_half_open_policy() {
584 let mut histogram =
585 JointHistogram::empty(vec![vec![0.0, 0.25, 2.0], vec![-1.0, 0.0, 0.5, 3.0]]).unwrap();
586 histogram.fill_weighted(&[0.0, -1.0], 1.0).unwrap();
587 histogram.fill_weighted(&[0.25, 0.0], 2.0).unwrap();
588
589 assert_eq!(histogram.shape(), [2, 3]);
590 assert_eq!(histogram.values(), [1.0, 0.0, 0.0, 0.0, 2.0, 0.0]);
591 }
592
593 #[test]
594 fn incompatible_merge_is_atomic() {
595 let mut left = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
596 left.fill_weighted(&[0.5], 2.0).unwrap();
597 let before = left.clone();
598 let other = JointHistogram::empty(vec![vec![0.0, 2.0]]).unwrap();
599
600 assert!(left.merge(&other).is_err());
601 assert_eq!(left, before);
602 }
603
604 #[test]
605 fn all_out_of_range_and_empty_histograms_remain_valid() {
606 let mut histogram = JointHistogram::empty(vec![vec![0.0, 1.0], vec![0.0, 1.0]]).unwrap();
607 assert_eq!(histogram.values(), [0.0]);
608 histogram.fill_weighted(&[-1.0, 0.5], 2.0).unwrap();
609 histogram.fill_weighted(&[0.5, 2.0], -3.0).unwrap();
610 assert_eq!(histogram.values(), [0.0]);
611 assert_eq!(histogram.diagnostics().out_of_range_count(), 2);
612 assert_eq!(histogram.diagnostics().out_of_range_weight(), -1.0);
613 }
614
615 #[test]
616 fn serialization_restores_fill_and_merge_capability() {
617 let mut histogram = JointHistogram::empty(vec![vec![0.0, 1.0]]).unwrap();
618 histogram.fill_weighted(&[0.5], 2.0).unwrap();
619 let json = serde_json::to_string(&histogram).unwrap();
620 let mut restored: JointHistogram = serde_json::from_str(&json).unwrap();
621 restored.merge(&histogram).unwrap();
622 assert_eq!(restored.values(), [4.0]);
623 assert_eq!(restored.sum_squared_weights(), [8.0]);
624 }
625
626 #[test]
627 fn shape_overflow_is_rejected_before_allocation() {
628 assert!(checked_bin_count(&[usize::MAX, 2]).is_err());
629 }
630}