1use serde::{Deserialize, Serialize};
9use thiserror::Error;
10
11pub const MAX_CONFORMAL_SAMPLE_COUNT: u64 = u64::MAX - 1;
17
18const MAX_EXACT_METRIC_COUNT: u64 = (1_u64 << 53) - 1;
19const MAX_SHORTEST_DECIMAL_COEFFICIENT: u128 = 99_999_999_999_999_999;
20const MAX_U128_POWER_OF_TEN: u32 = 38;
21
22#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
24#[serde(rename_all = "snake_case")]
25pub enum ConformalMultiTargetPolicy {
26 Marginal,
28 JointMax,
30}
31
32#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
34#[serde(rename_all = "snake_case")]
35pub enum ConformalSmallSamplePolicy {
36 Error,
38 Unbounded,
40}
41
42#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
44#[serde(rename_all = "snake_case", tag = "status", content = "value")]
45pub enum ConformalRadius {
46 Finite(f64),
47 Unbounded,
48}
49
50#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
52#[serde(deny_unknown_fields)]
53pub struct SplitConformalQuantile {
54 pub coverage: f64,
55 pub rank: u64,
58 pub radii: Vec<ConformalRadius>,
61}
62
63#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
68#[serde(rename_all = "snake_case", tag = "status")]
69pub enum RegressionIntervalCell {
70 Finite { lower: f64, upper: f64 },
71 Unbounded,
72}
73
74impl RegressionIntervalCell {
75 pub fn endpoints(self) -> (Option<f64>, Option<f64>) {
77 match self {
78 Self::Finite { lower, upper } => (Some(lower), Some(upper)),
79 Self::Unbounded => (None, None),
80 }
81 }
82
83 pub fn midpoint(self) -> Option<f64> {
85 match self {
86 Self::Finite { lower, upper } => Some(finite_midpoint(lower, upper)),
87 Self::Unbounded => None,
88 }
89 }
90}
91
92#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
94#[serde(deny_unknown_fields)]
95pub struct RegressionConformalInterval {
96 pub coverage: f64,
97 pub cells: Vec<Vec<RegressionIntervalCell>>,
98}
99
100#[derive(Clone, Copy, Debug, Eq, PartialEq)]
102pub enum ConformalMeasurementStatus {
103 Finite,
104 Unbounded,
105}
106
107#[derive(Clone, Debug, PartialEq)]
109pub struct RegressionConformalMetrics {
110 pub target_index: Option<usize>,
112 pub measurement_status: ConformalMeasurementStatus,
113 pub empirical_coverage: f64,
114 pub coverage_gap: f64,
115 pub mean_width: Option<f64>,
116 pub median_width: Option<f64>,
117 pub interval_score: Option<f64>,
118}
119
120#[derive(Clone, Debug, Eq, PartialEq, Error)]
122pub enum ConformalError {
123 #[error("conformal coverages must be non-empty")]
124 EmptyCoverages,
125
126 #[error("coverage at index {index} must be finite and strictly inside (0, 1)")]
127 InvalidCoverage { index: usize },
128
129 #[error("coverage at index {index} is not strictly greater than its predecessor")]
130 NonIncreasingCoverage { index: usize },
131
132 #[error("sample_count must be in 1..={MAX_CONFORMAL_SAMPLE_COUNT}, got {sample_count}")]
133 InvalidSampleCount { sample_count: u64 },
134
135 #[error("failed to recover a shortest decimal binary64 representation")]
136 DecimalConversion,
137
138 #[error("{matrix} matrix must be non-empty")]
139 EmptyMatrix { matrix: &'static str },
140
141 #[error("{matrix} row {row} must be non-empty")]
142 EmptyMatrixRow { matrix: &'static str, row: usize },
143
144 #[error("{matrix} row {row} has width {actual}, expected the rectangular width {expected}")]
145 RaggedMatrix {
146 matrix: &'static str,
147 row: usize,
148 expected: usize,
149 actual: usize,
150 },
151
152 #[error("{matrix} value at row {row}, target {target} must be finite")]
153 NonFiniteMatrixValue {
154 matrix: &'static str,
155 row: usize,
156 target: usize,
157 },
158
159 #[error("residual at row {row}, target {target} must be non-negative")]
160 NegativeResidual { row: usize, target: usize },
161
162 #[error("finite-sample rank {rank} exceeds calibration size {sample_count}")]
163 SmallSampleRank { rank: u64, sample_count: u64 },
164
165 #[error("split-conformal quantiles must be non-empty")]
166 EmptyQuantiles,
167
168 #[error("quantile {coverage_index} has rank zero")]
169 ZeroQuantileRank { coverage_index: usize },
170
171 #[error("quantile rank decreases at coverage index {coverage_index}")]
172 DecreasingQuantileRank { coverage_index: usize },
173
174 #[error(
175 "quantile {coverage_index} has {actual} radii, expected {expected} for this target policy"
176 )]
177 QuantileShape {
178 coverage_index: usize,
179 expected: usize,
180 actual: usize,
181 },
182
183 #[error("quantile {coverage_index}, radius {radius_index} must be finite and non-negative")]
184 InvalidRadius {
185 coverage_index: usize,
186 radius_index: usize,
187 },
188
189 #[error("quantile {coverage_index} mixes finite and unbounded radii")]
190 MixedRadiusStatus { coverage_index: usize },
191
192 #[error(
193 "quantile radius is not nested at coverage index {coverage_index}, radius {radius_index}"
194 )]
195 NonNestedRadius {
196 coverage_index: usize,
197 radius_index: usize,
198 },
199
200 #[error("{left} and {right} matrix row counts differ")]
201 MatrixRowCountMismatch {
202 left: &'static str,
203 right: &'static str,
204 },
205
206 #[error("{left} and {right} matrix target widths differ")]
207 MatrixTargetCountMismatch {
208 left: &'static str,
209 right: &'static str,
210 },
211
212 #[error("interval cell at row {row}, target {target} must contain ordered finite bounds")]
213 InvalidIntervalCell { row: usize, target: usize },
214
215 #[error("finite arithmetic overflow while computing {operation}")]
216 ArithmeticOverflow { operation: &'static str },
217
218 #[error(
219 "finite interval at coverage {coverage_index}, row {row}, target {target} cannot preserve the W0 decimal midpoint and radius closures"
220 )]
221 UnrepresentableInterval {
222 coverage_index: usize,
223 row: usize,
224 target: usize,
225 },
226
227 #[error("metric cell count exceeds the exact binary64 integer range")]
228 MetricCountTooLarge,
229}
230
231pub fn validate_conformal_coverages(coverages: &[f64]) -> Result<(), ConformalError> {
233 if coverages.is_empty() {
234 return Err(ConformalError::EmptyCoverages);
235 }
236 for (index, coverage) in coverages.iter().copied().enumerate() {
237 if !(coverage.is_finite() && 0.0 < coverage && coverage < 1.0) {
238 return Err(ConformalError::InvalidCoverage { index });
239 }
240 if index > 0 && coverages[index - 1] >= coverage {
241 return Err(ConformalError::NonIncreasingCoverage { index });
242 }
243 }
244 Ok(())
245}
246
247pub fn finite_sample_conformal_rank(
255 sample_count: u64,
256 coverage: f64,
257) -> Result<u64, ConformalError> {
258 if sample_count == 0 || sample_count > MAX_CONFORMAL_SAMPLE_COUNT {
259 return Err(ConformalError::InvalidSampleCount { sample_count });
260 }
261 validate_conformal_coverages(&[coverage])?;
262 let decimal = shortest_decimal(coverage)?;
263 let multiplier = u128::from(
264 sample_count
265 .checked_add(1)
266 .ok_or(ConformalError::InvalidSampleCount { sample_count })?,
267 );
268 let scaled =
269 multiplier
270 .checked_mul(decimal.coefficient)
271 .ok_or(ConformalError::ArithmeticOverflow {
272 operation: "finite-sample rank numerator",
273 })?;
274
275 let rank = if decimal.scale > MAX_U128_POWER_OF_TEN {
280 1_u128
281 } else {
282 let denominator = checked_power_of_ten(decimal.scale)?;
283 let quotient = scaled / denominator;
284 quotient + u128::from(scaled % denominator != 0)
285 };
286 u64::try_from(rank).map_err(|_| ConformalError::ArithmeticOverflow {
287 operation: "finite-sample rank result",
288 })
289}
290
291pub fn split_absolute_residual_quantiles(
296 residuals: &[Vec<f64>],
297 coverages: &[f64],
298 multi_target_policy: ConformalMultiTargetPolicy,
299 small_sample_policy: ConformalSmallSamplePolicy,
300) -> Result<Vec<SplitConformalQuantile>, ConformalError> {
301 let target_count = validate_finite_matrix(residuals, "residuals", true)?;
302 validate_conformal_coverages(coverages)?;
303 let sample_count =
304 u64::try_from(residuals.len()).map_err(|_| ConformalError::InvalidSampleCount {
305 sample_count: u64::MAX,
306 })?;
307 if sample_count > MAX_CONFORMAL_SAMPLE_COUNT {
308 return Err(ConformalError::InvalidSampleCount { sample_count });
309 }
310
311 let score_count = match multi_target_policy {
312 ConformalMultiTargetPolicy::Marginal => target_count,
313 ConformalMultiTargetPolicy::JointMax => 1,
314 };
315 let mut ordered_scores = (0..score_count)
316 .map(|_| Vec::with_capacity(residuals.len()))
317 .collect::<Vec<_>>();
318 for row in residuals {
319 match multi_target_policy {
320 ConformalMultiTargetPolicy::Marginal => {
321 for (target, residual) in row.iter().copied().enumerate() {
322 ordered_scores[target].push(normalized_zero(residual));
323 }
324 }
325 ConformalMultiTargetPolicy::JointMax => {
326 let maximum = row
327 .iter()
328 .copied()
329 .map(normalized_zero)
330 .fold(0.0_f64, f64::max);
331 ordered_scores[0].push(maximum);
332 }
333 }
334 }
335 for scores in &mut ordered_scores {
336 scores.sort_by(f64::total_cmp);
337 }
338
339 let mut quantiles = Vec::with_capacity(coverages.len());
340 for coverage in coverages.iter().copied() {
341 let rank = finite_sample_conformal_rank(sample_count, coverage)?;
342 let radii = if rank > sample_count {
343 match small_sample_policy {
344 ConformalSmallSamplePolicy::Error => {
345 return Err(ConformalError::SmallSampleRank { rank, sample_count });
346 }
347 ConformalSmallSamplePolicy::Unbounded => {
348 vec![ConformalRadius::Unbounded; score_count]
349 }
350 }
351 } else {
352 let index =
353 usize::try_from(rank - 1).map_err(|_| ConformalError::ArithmeticOverflow {
354 operation: "quantile order-statistic index",
355 })?;
356 ordered_scores
357 .iter()
358 .map(|scores| ConformalRadius::Finite(scores[index]))
359 .collect()
360 };
361 quantiles.push(SplitConformalQuantile {
362 coverage,
363 rank,
364 radii,
365 });
366 }
367 Ok(quantiles)
368}
369
370pub fn apply_split_absolute_residual(
380 point_predictions: &[Vec<f64>],
381 quantiles: &[SplitConformalQuantile],
382 multi_target_policy: ConformalMultiTargetPolicy,
383) -> Result<Vec<RegressionConformalInterval>, ConformalError> {
384 let target_count = validate_finite_matrix(point_predictions, "point predictions", false)?;
385 validate_quantiles(quantiles, target_count, multi_target_policy)?;
386
387 let mut intervals = Vec::with_capacity(quantiles.len());
388 for (coverage_index, quantile) in quantiles.iter().enumerate() {
389 let mut cells = Vec::with_capacity(point_predictions.len());
390 for (row_index, row) in point_predictions.iter().enumerate() {
391 let mut interval_row = Vec::with_capacity(target_count);
392 for (target, point) in row.iter().copied().enumerate() {
393 let radius = quantile.radii[match multi_target_policy {
394 ConformalMultiTargetPolicy::Marginal => target,
395 ConformalMultiTargetPolicy::JointMax => 0,
396 }];
397 let cell = match radius {
398 ConformalRadius::Unbounded => RegressionIntervalCell::Unbounded,
399 ConformalRadius::Finite(radius) => {
400 let radius = normalized_zero(radius);
401 let (lower, upper) = decimal_conformal_endpoints(point, radius)?;
402 if !decimal_interval_closes(point, radius, lower, upper)? {
403 return Err(ConformalError::UnrepresentableInterval {
404 coverage_index,
405 row: row_index,
406 target,
407 });
408 }
409 RegressionIntervalCell::Finite { lower, upper }
410 }
411 };
412 interval_row.push(cell);
413 }
414 cells.push(interval_row);
415 }
416 intervals.push(RegressionConformalInterval {
417 coverage: quantile.coverage,
418 cells,
419 });
420 }
421 Ok(intervals)
422}
423
424pub fn regression_conformal_metrics(
432 truth: &[Vec<f64>],
433 interval: &RegressionConformalInterval,
434 multi_target_policy: ConformalMultiTargetPolicy,
435) -> Result<Vec<RegressionConformalMetrics>, ConformalError> {
436 validate_conformal_coverages(&[interval.coverage])?;
437 let target_count = validate_finite_matrix(truth, "truth", false)?;
438 let interval_target_count = validate_interval_matrix(&interval.cells)?;
439 if truth.len() != interval.cells.len() {
440 return Err(ConformalError::MatrixRowCountMismatch {
441 left: "truth",
442 right: "interval",
443 });
444 }
445 if target_count != interval_target_count {
446 return Err(ConformalError::MatrixTargetCountMismatch {
447 left: "truth",
448 right: "interval",
449 });
450 }
451
452 let alpha = 1.0 - interval.coverage;
453 let miss_scale = 2.0 / alpha;
454 let mut covered = Vec::with_capacity(truth.len());
455 let mut widths = Vec::with_capacity(truth.len());
456 let mut scores = Vec::with_capacity(truth.len());
457 for (row_index, (truth_row, interval_row)) in truth.iter().zip(&interval.cells).enumerate() {
458 let mut covered_row = Vec::with_capacity(target_count);
459 let mut width_row = Vec::with_capacity(target_count);
460 let mut score_row = Vec::with_capacity(target_count);
461 for (target, (value, cell)) in truth_row
462 .iter()
463 .copied()
464 .zip(interval_row.iter().copied())
465 .enumerate()
466 {
467 match cell {
468 RegressionIntervalCell::Unbounded => {
469 covered_row.push(true);
470 width_row.push(None);
471 score_row.push(None);
472 }
473 RegressionIntervalCell::Finite { lower, upper } => {
474 if !lower.is_finite() || !upper.is_finite() || lower > upper {
475 return Err(ConformalError::InvalidIntervalCell {
476 row: row_index,
477 target,
478 });
479 }
480 let width = checked_finite(upper - lower, "interval width")?;
481 let miss_distance = if value < lower {
482 lower - value
483 } else if value > upper {
484 value - upper
485 } else {
486 0.0
487 };
488 let penalty = checked_finite(
489 miss_scale * checked_finite(miss_distance, "interval miss distance")?,
490 "Winkler miss penalty",
491 )?;
492 let score = checked_finite(width + penalty, "Winkler interval score")?;
493 covered_row.push(lower <= value && value <= upper);
494 width_row.push(Some(width));
495 score_row.push(Some(score));
496 }
497 }
498 }
499 covered.push(covered_row);
500 widths.push(width_row);
501 scores.push(score_row);
502 }
503
504 match multi_target_policy {
505 ConformalMultiTargetPolicy::Marginal => (0..target_count)
506 .map(|target| {
507 summarize_metrics(
508 interval.coverage,
509 covered.iter().map(|row| row[target]).collect(),
510 widths.iter().map(|row| row[target]).collect(),
511 scores.iter().map(|row| row[target]).collect(),
512 Some(target),
513 )
514 })
515 .collect(),
516 ConformalMultiTargetPolicy::JointMax => summarize_metrics(
517 interval.coverage,
518 covered
519 .iter()
520 .map(|row| row.iter().all(|value| *value))
521 .collect(),
522 widths.iter().flatten().copied().collect(),
523 scores.iter().flatten().copied().collect(),
524 None,
525 )
526 .map(|summary| vec![summary]),
527 }
528}
529
530#[derive(Clone, Copy, Debug)]
531struct ShortestDecimal {
532 coefficient: u128,
533 scale: u32,
534}
535
536#[derive(Clone, Debug, Eq, PartialEq)]
541struct ExactDecimal {
542 negative: bool,
543 digits: Vec<u8>,
544 exponent: i32,
545}
546
547impl ExactDecimal {
548 fn zero(negative: bool) -> Self {
549 Self {
550 negative,
551 digits: vec![0],
552 exponent: 0,
553 }
554 }
555
556 fn is_zero(&self) -> bool {
557 self.digits == [0]
558 }
559}
560
561fn shortest_decimal(value: f64) -> Result<ShortestDecimal, ConformalError> {
562 let rendered = value.to_string();
563 let (mantissa, exponent) = match rendered.find(['e', 'E']) {
564 Some(index) => {
565 let exponent = rendered[index + 1..]
566 .parse::<i32>()
567 .map_err(|_| ConformalError::DecimalConversion)?;
568 (&rendered[..index], exponent)
569 }
570 None => (rendered.as_str(), 0),
571 };
572 let mut coefficient = 0_u128;
573 let mut fractional_digits = 0_i32;
574 let mut after_decimal = false;
575 let mut saw_digit = false;
576 for byte in mantissa.bytes() {
577 match byte {
578 b'.' if !after_decimal => after_decimal = true,
579 b'0'..=b'9' => {
580 saw_digit = true;
581 coefficient = coefficient
582 .checked_mul(10)
583 .and_then(|current| current.checked_add(u128::from(byte - b'0')))
584 .ok_or(ConformalError::DecimalConversion)?;
585 if after_decimal {
586 fractional_digits = fractional_digits
587 .checked_add(1)
588 .ok_or(ConformalError::DecimalConversion)?;
589 }
590 }
591 _ => return Err(ConformalError::DecimalConversion),
592 }
593 }
594 if !saw_digit || coefficient == 0 || coefficient > MAX_SHORTEST_DECIMAL_COEFFICIENT {
595 return Err(ConformalError::DecimalConversion);
596 }
597 let scale = fractional_digits
598 .checked_sub(exponent)
599 .ok_or(ConformalError::DecimalConversion)?;
600 if scale <= 0 {
601 return Err(ConformalError::DecimalConversion);
602 }
603 Ok(ShortestDecimal {
604 coefficient,
605 scale: u32::try_from(scale).map_err(|_| ConformalError::DecimalConversion)?,
606 })
607}
608
609fn exact_decimal_from_f64(value: f64) -> Result<ExactDecimal, ConformalError> {
610 if !value.is_finite() {
611 return Err(ConformalError::DecimalConversion);
612 }
613 let rendered = value.to_string();
614 let (negative, unsigned) = rendered
615 .strip_prefix('-')
616 .map_or((false, rendered.as_str()), |rest| (true, rest));
617 let (mantissa, scientific_exponent) = match unsigned.find(['e', 'E']) {
618 Some(index) => {
619 let exponent = unsigned[index + 1..]
620 .parse::<i32>()
621 .map_err(|_| ConformalError::DecimalConversion)?;
622 (&unsigned[..index], exponent)
623 }
624 None => (unsigned, 0),
625 };
626
627 let mut digits = Vec::with_capacity(mantissa.len());
628 let mut fractional_digits = 0_i32;
629 let mut after_decimal = false;
630 for byte in mantissa.bytes() {
631 match byte {
632 b'.' if !after_decimal => after_decimal = true,
633 b'0'..=b'9' => {
634 digits.push(byte - b'0');
635 if after_decimal {
636 fractional_digits = fractional_digits
637 .checked_add(1)
638 .ok_or(ConformalError::DecimalConversion)?;
639 }
640 }
641 _ => return Err(ConformalError::DecimalConversion),
642 }
643 }
644 if digits.is_empty() {
645 return Err(ConformalError::DecimalConversion);
646 }
647
648 let first_nonzero = digits.iter().position(|digit| *digit != 0);
649 let Some(first_nonzero) = first_nonzero else {
650 return Ok(ExactDecimal::zero(negative));
651 };
652 digits.drain(..first_nonzero);
653 let exponent = scientific_exponent
654 .checked_sub(fractional_digits)
655 .ok_or(ConformalError::DecimalConversion)?;
656 normalize_exact_decimal(ExactDecimal {
657 negative,
658 digits,
659 exponent,
660 })
661}
662
663fn normalize_exact_decimal(mut value: ExactDecimal) -> Result<ExactDecimal, ConformalError> {
664 while value.digits.len() > 1 && value.digits.last() == Some(&0) {
665 value.digits.pop();
666 value.exponent = value
667 .exponent
668 .checked_add(1)
669 .ok_or(ConformalError::DecimalConversion)?;
670 }
671 Ok(value)
672}
673
674fn aligned_decimal_digits(
675 value: &ExactDecimal,
676 common_exponent: i32,
677) -> Result<Vec<u8>, ConformalError> {
678 let trailing_zeros = value
679 .exponent
680 .checked_sub(common_exponent)
681 .and_then(|count| usize::try_from(count).ok())
682 .ok_or(ConformalError::DecimalConversion)?;
683 let expanded_len = value
684 .digits
685 .len()
686 .checked_add(trailing_zeros)
687 .ok_or(ConformalError::DecimalConversion)?;
688 let mut digits = Vec::with_capacity(expanded_len);
689 digits.extend_from_slice(&value.digits);
690 digits.resize(expanded_len, 0);
691 Ok(digits)
692}
693
694fn add_decimal_magnitudes(left: &[u8], right: &[u8]) -> Vec<u8> {
695 let digit_count = left.len().max(right.len());
696 let mut left = left.iter().rev();
697 let mut right = right.iter().rev();
698 let mut reversed = Vec::with_capacity(digit_count + 1);
699 let mut carry = 0_u8;
700 for _ in 0..digit_count {
701 let sum = left.next().copied().unwrap_or(0) + right.next().copied().unwrap_or(0) + carry;
702 reversed.push(sum % 10);
703 carry = sum / 10;
704 }
705 if carry != 0 {
706 reversed.push(carry);
707 }
708 reversed.reverse();
709 reversed
710}
711
712fn subtract_decimal_magnitudes(larger: &[u8], smaller: &[u8]) -> Result<Vec<u8>, ConformalError> {
713 let mut smaller = smaller.iter().rev();
714 let mut reversed = Vec::with_capacity(larger.len());
715 let mut borrow = 0_i16;
716 for larger_digit in larger.iter().rev().copied() {
717 let smaller_digit = i16::from(smaller.next().copied().unwrap_or(0));
718 let mut difference = i16::from(larger_digit) - smaller_digit - borrow;
719 if difference < 0 {
720 difference += 10;
721 borrow = 1;
722 } else {
723 borrow = 0;
724 }
725 reversed.push(u8::try_from(difference).map_err(|_| ConformalError::DecimalConversion)?);
726 }
727 if borrow != 0 {
728 return Err(ConformalError::DecimalConversion);
729 }
730 reversed.reverse();
731 let leading_zeros = reversed
732 .iter()
733 .position(|digit| *digit != 0)
734 .unwrap_or(reversed.len().saturating_sub(1));
735 reversed.drain(..leading_zeros);
736 Ok(reversed)
737}
738
739fn add_exact_decimals(
740 left: &ExactDecimal,
741 right: &ExactDecimal,
742 subtract_right: bool,
743) -> Result<ExactDecimal, ConformalError> {
744 if left.is_zero() && right.is_zero() {
745 return Ok(ExactDecimal::zero(false));
746 }
747 if right.is_zero() {
748 return Ok(left.clone());
749 }
750 if left.is_zero() {
751 let mut result = right.clone();
752 result.negative ^= subtract_right;
753 return Ok(result);
754 }
755
756 let common_exponent = left.exponent.min(right.exponent);
757 let left_digits = aligned_decimal_digits(left, common_exponent)?;
758 let right_digits = aligned_decimal_digits(right, common_exponent)?;
759 let right_negative = right.negative ^ subtract_right;
760 let (negative, digits) = if left.negative == right_negative {
761 (
762 left.negative,
763 add_decimal_magnitudes(&left_digits, &right_digits),
764 )
765 } else {
766 let magnitude_order = left_digits
767 .len()
768 .cmp(&right_digits.len())
769 .then_with(|| left_digits.cmp(&right_digits));
770 match magnitude_order {
771 std::cmp::Ordering::Greater => (
772 left.negative,
773 subtract_decimal_magnitudes(&left_digits, &right_digits)?,
774 ),
775 std::cmp::Ordering::Less => (
776 right_negative,
777 subtract_decimal_magnitudes(&right_digits, &left_digits)?,
778 ),
779 std::cmp::Ordering::Equal => return Ok(ExactDecimal::zero(false)),
780 }
781 };
782 normalize_exact_decimal(ExactDecimal {
783 negative,
784 digits,
785 exponent: common_exponent,
786 })
787}
788
789fn exact_decimal_to_f64(
790 value: &ExactDecimal,
791 operation: &'static str,
792) -> Result<f64, ConformalError> {
793 if value.is_zero() {
794 return Ok(f64::from_bits(u64::from(value.negative) << 63));
795 }
796 let scientific_exponent = value
797 .exponent
798 .checked_add(
799 i32::try_from(value.digits.len() - 1).map_err(|_| ConformalError::DecimalConversion)?,
800 )
801 .ok_or(ConformalError::DecimalConversion)?;
802 let mut rendered = String::with_capacity(value.digits.len() + 16);
803 if value.negative {
804 rendered.push('-');
805 }
806 rendered.push(char::from(b'0' + value.digits[0]));
807 if value.digits.len() > 1 {
808 rendered.push('.');
809 rendered.extend(
810 value.digits[1..]
811 .iter()
812 .map(|digit| char::from(b'0' + digit)),
813 );
814 }
815 rendered.push('e');
816 rendered.push_str(&scientific_exponent.to_string());
817 let rounded = rendered
818 .parse::<f64>()
819 .map_err(|_| ConformalError::DecimalConversion)?;
820 checked_finite(rounded, operation)
821}
822
823fn decimal_conformal_endpoints(point: f64, radius: f64) -> Result<(f64, f64), ConformalError> {
824 let point_decimal = exact_decimal_from_f64(point)?;
825 let radius_decimal = exact_decimal_from_f64(radius)?;
826 let mut lower_decimal = add_exact_decimals(&point_decimal, &radius_decimal, true)?;
827 let mut upper_decimal = add_exact_decimals(&point_decimal, &radius_decimal, false)?;
828
829 if lower_decimal.is_zero() {
832 lower_decimal.negative = point == 0.0 && point.is_sign_negative();
833 }
834 if upper_decimal.is_zero() {
835 upper_decimal.negative = false;
836 }
837
838 Ok((
839 exact_decimal_to_f64(&lower_decimal, "finite conformal lower endpoint")?,
840 exact_decimal_to_f64(&upper_decimal, "finite conformal upper endpoint")?,
841 ))
842}
843
844fn exact_decimal_values_equal(left: &ExactDecimal, right: &ExactDecimal) -> bool {
845 (left.is_zero() && right.is_zero()) || left == right
846}
847
848fn decimal_interval_closes(
849 point: f64,
850 radius: f64,
851 lower: f64,
852 upper: f64,
853) -> Result<bool, ConformalError> {
854 let point = exact_decimal_from_f64(point)?;
855 let radius = exact_decimal_from_f64(radius)?;
856 let lower = exact_decimal_from_f64(lower)?;
857 let upper = exact_decimal_from_f64(upper)?;
858 let endpoint_midpoint = add_exact_decimals(&lower, &upper, false)?;
859 let expected_midpoint = add_exact_decimals(&point, &point, false)?;
860 let endpoint_width = add_exact_decimals(&upper, &lower, true)?;
861 let expected_width = add_exact_decimals(&radius, &radius, false)?;
862 Ok(
863 exact_decimal_values_equal(&endpoint_midpoint, &expected_midpoint)
864 && exact_decimal_values_equal(&endpoint_width, &expected_width),
865 )
866}
867
868fn checked_power_of_ten(power: u32) -> Result<u128, ConformalError> {
869 let mut value = 1_u128;
870 for _ in 0..power {
871 value = value
872 .checked_mul(10)
873 .ok_or(ConformalError::ArithmeticOverflow {
874 operation: "decimal rank denominator",
875 })?;
876 }
877 Ok(value)
878}
879
880fn validate_finite_matrix(
881 matrix: &[Vec<f64>],
882 name: &'static str,
883 non_negative: bool,
884) -> Result<usize, ConformalError> {
885 if matrix.is_empty() {
886 return Err(ConformalError::EmptyMatrix { matrix: name });
887 }
888 let expected = matrix[0].len();
889 if expected == 0 {
890 return Err(ConformalError::EmptyMatrixRow {
891 matrix: name,
892 row: 0,
893 });
894 }
895 for (row_index, row) in matrix.iter().enumerate() {
896 if row.len() != expected {
897 return Err(ConformalError::RaggedMatrix {
898 matrix: name,
899 row: row_index,
900 expected,
901 actual: row.len(),
902 });
903 }
904 for (target, value) in row.iter().copied().enumerate() {
905 if !value.is_finite() {
906 return Err(ConformalError::NonFiniteMatrixValue {
907 matrix: name,
908 row: row_index,
909 target,
910 });
911 }
912 if non_negative && value < 0.0 {
913 return Err(ConformalError::NegativeResidual {
914 row: row_index,
915 target,
916 });
917 }
918 }
919 }
920 Ok(expected)
921}
922
923fn validate_interval_matrix(
924 matrix: &[Vec<RegressionIntervalCell>],
925) -> Result<usize, ConformalError> {
926 if matrix.is_empty() {
927 return Err(ConformalError::EmptyMatrix { matrix: "interval" });
928 }
929 let expected = matrix[0].len();
930 if expected == 0 {
931 return Err(ConformalError::EmptyMatrixRow {
932 matrix: "interval",
933 row: 0,
934 });
935 }
936 for (row, cells) in matrix.iter().enumerate() {
937 if cells.len() != expected {
938 return Err(ConformalError::RaggedMatrix {
939 matrix: "interval",
940 row,
941 expected,
942 actual: cells.len(),
943 });
944 }
945 for (target, cell) in cells.iter().copied().enumerate() {
946 if let RegressionIntervalCell::Finite { lower, upper } = cell {
947 if !lower.is_finite() || !upper.is_finite() || lower > upper {
948 return Err(ConformalError::InvalidIntervalCell { row, target });
949 }
950 }
951 }
952 }
953 Ok(expected)
954}
955
956fn validate_quantiles(
957 quantiles: &[SplitConformalQuantile],
958 target_count: usize,
959 policy: ConformalMultiTargetPolicy,
960) -> Result<(), ConformalError> {
961 if quantiles.is_empty() {
962 return Err(ConformalError::EmptyQuantiles);
963 }
964 let coverages = quantiles
965 .iter()
966 .map(|quantile| quantile.coverage)
967 .collect::<Vec<_>>();
968 validate_conformal_coverages(&coverages)?;
969 let expected = match policy {
970 ConformalMultiTargetPolicy::Marginal => target_count,
971 ConformalMultiTargetPolicy::JointMax => 1,
972 };
973 let mut previous_rank = 0_u64;
974 let mut previous = vec![None; expected];
975 for (coverage_index, quantile) in quantiles.iter().enumerate() {
976 if quantile.rank == 0 {
977 return Err(ConformalError::ZeroQuantileRank { coverage_index });
978 }
979 if coverage_index > 0 && quantile.rank < previous_rank {
980 return Err(ConformalError::DecreasingQuantileRank { coverage_index });
981 }
982 previous_rank = quantile.rank;
983 if quantile.radii.len() != expected {
984 return Err(ConformalError::QuantileShape {
985 coverage_index,
986 expected,
987 actual: quantile.radii.len(),
988 });
989 }
990 let first_is_unbounded = matches!(quantile.radii[0], ConformalRadius::Unbounded);
991 if quantile
992 .radii
993 .iter()
994 .any(|radius| matches!(radius, ConformalRadius::Unbounded) != first_is_unbounded)
995 {
996 return Err(ConformalError::MixedRadiusStatus { coverage_index });
997 }
998 for (radius_index, radius) in quantile.radii.iter().copied().enumerate() {
999 if let ConformalRadius::Finite(value) = radius {
1000 if !value.is_finite() || value < 0.0 {
1001 return Err(ConformalError::InvalidRadius {
1002 coverage_index,
1003 radius_index,
1004 });
1005 }
1006 }
1007 if let Some(previous_radius) = previous[radius_index] {
1008 let nested = match (previous_radius, radius) {
1009 (ConformalRadius::Finite(left), ConformalRadius::Finite(right)) => {
1010 left <= right
1011 }
1012 (ConformalRadius::Finite(_), ConformalRadius::Unbounded)
1013 | (ConformalRadius::Unbounded, ConformalRadius::Unbounded) => true,
1014 (ConformalRadius::Unbounded, ConformalRadius::Finite(_)) => false,
1015 };
1016 if !nested {
1017 return Err(ConformalError::NonNestedRadius {
1018 coverage_index,
1019 radius_index,
1020 });
1021 }
1022 }
1023 previous[radius_index] = Some(radius);
1024 }
1025 }
1026 Ok(())
1027}
1028
1029fn summarize_metrics(
1030 coverage: f64,
1031 covered: Vec<bool>,
1032 widths: Vec<Option<f64>>,
1033 scores: Vec<Option<f64>>,
1034 target_index: Option<usize>,
1035) -> Result<RegressionConformalMetrics, ConformalError> {
1036 let count = u64::try_from(covered.len()).map_err(|_| ConformalError::MetricCountTooLarge)?;
1037 if count == 0 || count > MAX_EXACT_METRIC_COUNT {
1038 return Err(ConformalError::MetricCountTooLarge);
1039 }
1040 let measurement_count =
1041 u64::try_from(widths.len()).map_err(|_| ConformalError::MetricCountTooLarge)?;
1042 if measurement_count == 0
1043 || measurement_count > MAX_EXACT_METRIC_COUNT
1044 || widths.len() != scores.len()
1045 {
1046 return Err(ConformalError::MetricCountTooLarge);
1047 }
1048 let covered_count = u64::try_from(covered.iter().filter(|value| **value).count())
1049 .map_err(|_| ConformalError::MetricCountTooLarge)?;
1050 let empirical_coverage = (covered_count as f64) / (count as f64);
1051 let coverage_gap = empirical_coverage - coverage;
1052 if widths.iter().chain(&scores).any(Option::is_none) {
1053 return Ok(RegressionConformalMetrics {
1054 target_index,
1055 measurement_status: ConformalMeasurementStatus::Unbounded,
1056 empirical_coverage,
1057 coverage_gap,
1058 mean_width: None,
1059 median_width: None,
1060 interval_score: None,
1061 });
1062 }
1063
1064 let mut finite_widths = widths.into_iter().flatten().collect::<Vec<_>>();
1065 let finite_scores = scores.into_iter().flatten().collect::<Vec<_>>();
1066 finite_widths.sort_by(f64::total_cmp);
1067 let mean_width = checked_mean(&finite_widths, "mean interval width")?;
1068 let interval_score = checked_mean(&finite_scores, "mean Winkler interval score")?;
1069 let middle = finite_widths.len() / 2;
1070 let median_width = if finite_widths.len() % 2 == 1 {
1071 finite_widths[middle]
1072 } else {
1073 finite_midpoint(finite_widths[middle - 1], finite_widths[middle])
1074 };
1075 Ok(RegressionConformalMetrics {
1076 target_index,
1077 measurement_status: ConformalMeasurementStatus::Finite,
1078 empirical_coverage,
1079 coverage_gap,
1080 mean_width: Some(mean_width),
1081 median_width: Some(median_width),
1082 interval_score: Some(interval_score),
1083 })
1084}
1085
1086fn checked_mean(values: &[f64], operation: &'static str) -> Result<f64, ConformalError> {
1087 if values.is_empty() {
1088 return Err(ConformalError::MetricCountTooLarge);
1089 }
1090 let count = u64::try_from(values.len()).map_err(|_| ConformalError::MetricCountTooLarge)?;
1091 let mut sum = 0.0_f64;
1092 let mut sum_is_finite = true;
1093 for value in values.iter().copied() {
1094 let next = sum + value;
1095 if !next.is_finite() {
1096 sum_is_finite = false;
1097 break;
1098 }
1099 sum = next;
1100 }
1101 if sum_is_finite {
1102 return Ok(sum / (count as f64));
1106 }
1107
1108 let mut mean = 0.0_f64;
1109 for (index, value) in values.iter().copied().enumerate() {
1110 let count = u64::try_from(index + 1).map_err(|_| ConformalError::MetricCountTooLarge)?;
1111 let delta = checked_finite(value - mean, operation)?;
1112 mean = checked_finite(mean + (delta / (count as f64)), operation)?;
1113 }
1114 Ok(mean)
1115}
1116
1117fn checked_finite(value: f64, operation: &'static str) -> Result<f64, ConformalError> {
1118 if value.is_finite() {
1119 Ok(value)
1120 } else {
1121 Err(ConformalError::ArithmeticOverflow { operation })
1122 }
1123}
1124
1125fn normalized_zero(value: f64) -> f64 {
1126 if value == 0.0 {
1127 0.0
1128 } else {
1129 value
1130 }
1131}
1132
1133fn finite_midpoint(lower: f64, upper: f64) -> f64 {
1134 let opposite_signs = lower.is_sign_negative() != upper.is_sign_negative();
1135 if opposite_signs || (lower.abs() <= f64::MAX / 2.0 && upper.abs() <= f64::MAX / 2.0) {
1136 (lower + upper) / 2.0
1137 } else {
1138 (lower / 2.0) + (upper / 2.0)
1139 }
1140}
1141
1142#[cfg(test)]
1143mod tests {
1144 use super::*;
1145
1146 fn assert_close(actual: f64, expected: f64) {
1147 assert!(
1148 (actual - expected).abs() <= 1e-12,
1149 "expected {expected:?}, got {actual:?}"
1150 );
1151 }
1152
1153 fn finite(value: f64) -> ConformalRadius {
1154 ConformalRadius::Finite(value)
1155 }
1156
1157 fn finite_cell(lower: f64, upper: f64) -> RegressionIntervalCell {
1158 RegressionIntervalCell::Finite { lower, upper }
1159 }
1160
1161 #[test]
1162 fn exact_rank_matches_frozen_standard_coverages() {
1163 let cases = [(0.8, 17), (0.9, 19), (0.95, 20), (0.99, 21), (0.999, 21)];
1164 for (coverage, rank) in cases {
1165 assert_eq!(finite_sample_conformal_rank(20, coverage).unwrap(), rank);
1166 }
1167 }
1168
1169 #[test]
1170 fn exact_rank_avoids_naive_binary64_boundary_drift() {
1171 assert_eq!((25.0_f64 * 0.28).ceil() as u64, 8);
1172 assert_eq!(finite_sample_conformal_rank(24, 0.28).unwrap(), 7);
1173
1174 let above_one_third = 0.333_333_333_333_333_37_f64;
1175 assert_eq!((3.0 * above_one_third).ceil() as u64, 1);
1176 assert_eq!(finite_sample_conformal_rank(2, above_one_third).unwrap(), 2);
1177 }
1178
1179 #[test]
1180 fn exact_rank_handles_binary64_extremes_and_sample_limit() {
1181 let minimum_subnormal = f64::from_bits(1);
1182 assert_eq!(
1183 finite_sample_conformal_rank(MAX_CONFORMAL_SAMPLE_COUNT, minimum_subnormal).unwrap(),
1184 1
1185 );
1186 assert_eq!(
1187 finite_sample_conformal_rank(MAX_CONFORMAL_SAMPLE_COUNT, f64::MIN_POSITIVE).unwrap(),
1188 1
1189 );
1190 assert_eq!(
1191 finite_sample_conformal_rank(1, f64::from_bits(1.0_f64.to_bits() - 1)).unwrap(),
1192 2
1193 );
1194 assert!(matches!(
1195 finite_sample_conformal_rank(u64::MAX, 0.5),
1196 Err(ConformalError::InvalidSampleCount { .. })
1197 ));
1198 assert_eq!(
1199 finite_sample_conformal_rank(MAX_CONFORMAL_SAMPLE_COUNT, 0.999_999_999_999_999_9)
1200 .unwrap(),
1201 18_446_744_073_709_549_771
1202 );
1203 }
1204
1205 #[test]
1206 fn exact_rank_exercises_the_u128_power_of_ten_boundary() {
1207 assert_eq!(checked_power_of_ten(38).unwrap(), 10_u128.pow(38));
1208 assert_eq!(finite_sample_conformal_rank(1, 1.0e-38).unwrap(), 1);
1209 }
1210
1211 #[test]
1212 fn calibration_covers_odd_even_ties_and_small_n() {
1213 let odd = vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0], vec![5.0]];
1214 let even = vec![vec![1.0], vec![2.0], vec![3.0], vec![4.0]];
1215 let ties = vec![vec![1.0], vec![2.0], vec![2.0], vec![4.0]];
1216 for (residuals, expected) in [(&odd, 3.0), (&even, 3.0), (&ties, 2.0)] {
1217 let result = split_absolute_residual_quantiles(
1218 residuals,
1219 &[0.5],
1220 ConformalMultiTargetPolicy::Marginal,
1221 ConformalSmallSamplePolicy::Error,
1222 )
1223 .unwrap();
1224 assert_eq!(result[0].rank, 3);
1225 assert_eq!(result[0].radii, vec![finite(expected)]);
1226 }
1227
1228 let one = vec![vec![1.0, 2.0]];
1229 assert!(matches!(
1230 split_absolute_residual_quantiles(
1231 &one,
1232 &[0.9],
1233 ConformalMultiTargetPolicy::Marginal,
1234 ConformalSmallSamplePolicy::Error,
1235 ),
1236 Err(ConformalError::SmallSampleRank {
1237 rank: 2,
1238 sample_count: 1
1239 })
1240 ));
1241 let unbounded = split_absolute_residual_quantiles(
1242 &one,
1243 &[0.9],
1244 ConformalMultiTargetPolicy::Marginal,
1245 ConformalSmallSamplePolicy::Unbounded,
1246 )
1247 .unwrap();
1248 assert_eq!(
1249 unbounded[0].radii,
1250 vec![ConformalRadius::Unbounded, ConformalRadius::Unbounded]
1251 );
1252 }
1253
1254 #[test]
1255 fn frozen_w0_unsorted_residual_quantiles_match() {
1256 let residuals = [
1257 0.4, 1.0, 0.2, 1.8, 0.6, 1.2, 0.8, 2.0, 0.1, 1.6, 0.3, 1.4, 0.5, 1.9, 0.7, 1.1, 0.9,
1258 1.3, 1.5, 1.7,
1259 ]
1260 .into_iter()
1261 .map(|value| vec![value])
1262 .collect::<Vec<_>>();
1263 let quantiles = split_absolute_residual_quantiles(
1264 &residuals,
1265 &[0.9, 0.95],
1266 ConformalMultiTargetPolicy::Marginal,
1267 ConformalSmallSamplePolicy::Error,
1268 )
1269 .unwrap();
1270 assert_eq!(quantiles[0].rank, 19);
1271 assert_eq!(quantiles[0].radii, vec![finite(1.9)]);
1272 assert_eq!(quantiles[1].rank, 20);
1273 assert_eq!(quantiles[1].radii, vec![finite(2.0)]);
1274 }
1275
1276 #[test]
1277 fn marginal_and_joint_max_calibration_are_distinct() {
1278 let residuals = vec![
1279 vec![1.0, 4.0],
1280 vec![2.0, 1.0],
1281 vec![3.0, 3.0],
1282 vec![4.0, 2.0],
1283 ];
1284 let marginal = split_absolute_residual_quantiles(
1285 &residuals,
1286 &[0.5],
1287 ConformalMultiTargetPolicy::Marginal,
1288 ConformalSmallSamplePolicy::Error,
1289 )
1290 .unwrap();
1291 assert_eq!(marginal[0].radii, vec![finite(3.0), finite(3.0)]);
1292
1293 let joint = split_absolute_residual_quantiles(
1294 &residuals,
1295 &[0.5],
1296 ConformalMultiTargetPolicy::JointMax,
1297 ConformalSmallSamplePolicy::Error,
1298 )
1299 .unwrap();
1300 assert_eq!(joint[0].radii, vec![finite(4.0)]);
1301 }
1302
1303 #[test]
1304 fn multi_coverage_application_is_nested_and_preserves_midpoints() {
1305 let residuals = (1..=20)
1306 .map(|value| vec![f64::from(value), f64::from(value) / 2.0])
1307 .collect::<Vec<_>>();
1308 let quantiles = split_absolute_residual_quantiles(
1309 &residuals,
1310 &[0.8, 0.9, 0.95, 0.99],
1311 ConformalMultiTargetPolicy::Marginal,
1312 ConformalSmallSamplePolicy::Unbounded,
1313 )
1314 .unwrap();
1315 let points = vec![vec![100.0, 10.0], vec![200.0, 20.0]];
1316 let intervals = apply_split_absolute_residual(
1317 &points,
1318 &quantiles,
1319 ConformalMultiTargetPolicy::Marginal,
1320 )
1321 .unwrap();
1322 assert_eq!(intervals.len(), 4);
1323 assert_eq!(intervals[0].cells[0][0], finite_cell(83.0, 117.0));
1324 assert_eq!(intervals[1].cells[0][0], finite_cell(81.0, 119.0));
1325 assert_eq!(intervals[2].cells[0][0], finite_cell(80.0, 120.0));
1326 assert_eq!(intervals[3].cells[0][0], RegressionIntervalCell::Unbounded);
1327 assert_eq!(intervals[0].cells[0][0].midpoint(), Some(100.0));
1328 assert_eq!(intervals[3].cells[0][0].endpoints(), (None, None));
1329 }
1330
1331 #[test]
1332 fn joint_radius_expands_to_every_prediction_target() {
1333 let quantiles = vec![SplitConformalQuantile {
1334 coverage: 0.8,
1335 rank: 4,
1336 radii: vec![finite(2.0)],
1337 }];
1338 let intervals = apply_split_absolute_residual(
1339 &[vec![10.0, 20.0]],
1340 &quantiles,
1341 ConformalMultiTargetPolicy::JointMax,
1342 )
1343 .unwrap();
1344 assert_eq!(
1345 intervals[0].cells[0],
1346 vec![finite_cell(8.0, 12.0), finite_cell(18.0, 22.0)]
1347 );
1348 }
1349
1350 #[test]
1351 fn decimal_endpoints_match_w0_instead_of_binary64_intermediate_arithmetic() {
1352 let quantiles = vec![SplitConformalQuantile {
1353 coverage: 0.8,
1354 rank: 1,
1355 radii: vec![finite(0.2)],
1356 }];
1357 let interval = apply_split_absolute_residual(
1358 &[vec![0.1]],
1359 &quantiles,
1360 ConformalMultiTargetPolicy::Marginal,
1361 )
1362 .unwrap();
1363 let RegressionIntervalCell::Finite { lower, upper } = interval[0].cells[0][0] else {
1364 panic!("W0 decimal endpoints must be finite");
1365 };
1366 assert_eq!(lower.to_bits(), (-0.1_f64).to_bits());
1367 assert_eq!(upper.to_bits(), 0.3_f64.to_bits());
1368 assert_ne!(upper.to_bits(), (0.1_f64 + 0.2_f64).to_bits());
1369 assert!(decimal_interval_closes(0.1, 0.2, lower, upper).unwrap());
1370
1371 let (lower, upper) = decimal_conformal_endpoints(100.0, 90.0).unwrap();
1372 assert_eq!((lower, upper), (10.0, 190.0));
1373 assert!(decimal_interval_closes(100.0, 90.0, lower, upper).unwrap());
1374 let (lower, upper) = decimal_conformal_endpoints(-100.0, 90.0).unwrap();
1375 assert_eq!((lower, upper), (-190.0, -10.0));
1376 assert!(decimal_interval_closes(-100.0, 90.0, lower, upper).unwrap());
1377 }
1378
1379 #[test]
1380 fn decimal_endpoints_reject_a_radius_below_the_point_ulp() {
1381 let quantiles = vec![SplitConformalQuantile {
1382 coverage: 0.8,
1383 rank: 1,
1384 radii: vec![finite(0.5)],
1385 }];
1386 assert!(matches!(
1387 apply_split_absolute_residual(
1388 &[vec![1.0e16]],
1389 &quantiles,
1390 ConformalMultiTargetPolicy::Marginal,
1391 ),
1392 Err(ConformalError::UnrepresentableInterval {
1393 coverage_index: 0,
1394 row: 0,
1395 target: 0,
1396 })
1397 ));
1398 }
1399
1400 #[test]
1401 fn decimal_endpoints_preserve_signed_zero_and_minimum_subnormal() {
1402 let zero_quantile = vec![SplitConformalQuantile {
1403 coverage: 0.5,
1404 rank: 1,
1405 radii: vec![finite(-0.0)],
1406 }];
1407 let zero_interval = apply_split_absolute_residual(
1408 &[vec![-0.0]],
1409 &zero_quantile,
1410 ConformalMultiTargetPolicy::Marginal,
1411 )
1412 .unwrap();
1413 let RegressionIntervalCell::Finite { lower, upper } = zero_interval[0].cells[0][0] else {
1414 panic!("zero-radius interval must be finite");
1415 };
1416 assert_eq!(lower.to_bits(), (-0.0_f64).to_bits());
1417 assert_eq!(upper.to_bits(), 0.0_f64.to_bits());
1418 assert!(decimal_interval_closes(-0.0, 0.0, lower, upper).unwrap());
1419
1420 let minimum_subnormal = f64::from_bits(1);
1421 let subnormal_quantile = vec![SplitConformalQuantile {
1422 coverage: 0.5,
1423 rank: 1,
1424 radii: vec![finite(minimum_subnormal)],
1425 }];
1426 let subnormal_interval = apply_split_absolute_residual(
1427 &[vec![0.0], vec![minimum_subnormal]],
1428 &subnormal_quantile,
1429 ConformalMultiTargetPolicy::Marginal,
1430 )
1431 .unwrap();
1432 assert_eq!(
1433 subnormal_interval[0].cells[0][0],
1434 finite_cell(-minimum_subnormal, minimum_subnormal)
1435 );
1436 assert_eq!(
1437 subnormal_interval[0].cells[1][0],
1438 finite_cell(0.0, f64::from_bits(2))
1439 );
1440 for (point, cell) in [0.0, minimum_subnormal]
1441 .into_iter()
1442 .zip(&subnormal_interval[0].cells)
1443 {
1444 let RegressionIntervalCell::Finite { lower, upper } = cell[0] else {
1445 panic!("subnormal interval must be finite");
1446 };
1447 assert!(decimal_interval_closes(point, minimum_subnormal, lower, upper).unwrap());
1448 }
1449 }
1450
1451 #[test]
1452 fn midpoint_is_stable_for_subnormal_and_extreme_bounds() {
1453 let minimum_subnormal = f64::from_bits(1);
1454 assert_eq!(
1455 finite_cell(minimum_subnormal, minimum_subnormal).midpoint(),
1456 Some(minimum_subnormal)
1457 );
1458 assert_eq!(finite_cell(f64::MAX, f64::MAX).midpoint(), Some(f64::MAX));
1459 assert_eq!(finite_cell(-f64::MAX, f64::MAX).midpoint(), Some(0.0));
1460 assert_eq!(
1461 finite_cell(-minimum_subnormal, 0.0)
1462 .midpoint()
1463 .unwrap()
1464 .to_bits(),
1465 (-0.0_f64).to_bits()
1466 );
1467 assert_eq!(
1468 finite_cell(minimum_subnormal, f64::from_bits(2)).midpoint(),
1469 Some(f64::from_bits(2))
1470 );
1471 }
1472
1473 #[test]
1474 fn finite_metrics_match_marginal_and_joint_w0_semantics() {
1475 let truth = vec![vec![1.0, 10.0], vec![3.0, 20.0]];
1476 let interval = RegressionConformalInterval {
1477 coverage: 0.8,
1478 cells: vec![
1479 vec![finite_cell(0.0, 2.0), finite_cell(9.0, 11.0)],
1480 vec![finite_cell(0.0, 2.0), finite_cell(19.0, 21.0)],
1481 ],
1482 };
1483 let marginal =
1484 regression_conformal_metrics(&truth, &interval, ConformalMultiTargetPolicy::Marginal)
1485 .unwrap();
1486 assert_eq!(marginal.len(), 2);
1487 assert_close(marginal[0].empirical_coverage, 0.5);
1488 assert_close(marginal[0].coverage_gap, -0.3);
1489 assert_close(marginal[0].mean_width.unwrap(), 2.0);
1490 assert_close(marginal[0].median_width.unwrap(), 2.0);
1491 assert_close(marginal[0].interval_score.unwrap(), 7.0);
1492 assert_close(marginal[1].empirical_coverage, 1.0);
1493 assert_close(marginal[1].interval_score.unwrap(), 2.0);
1494
1495 let joint =
1496 regression_conformal_metrics(&truth, &interval, ConformalMultiTargetPolicy::JointMax)
1497 .unwrap();
1498 assert_eq!(joint.len(), 1);
1499 assert_eq!(joint[0].target_index, None);
1500 assert_close(joint[0].empirical_coverage, 0.5);
1501 assert_close(joint[0].mean_width.unwrap(), 2.0);
1502 assert_close(joint[0].median_width.unwrap(), 2.0);
1503 assert_close(joint[0].interval_score.unwrap(), 4.5);
1504 }
1505
1506 #[test]
1507 fn frozen_w0_prediction_blocks_and_metrics_match() {
1508 let points = vec![vec![10.0, 20.0], vec![11.0, 21.0]];
1509 let truth = vec![vec![10.0, 21.0], vec![14.0, 21.0]];
1510
1511 let marginal_quantiles = vec![
1512 SplitConformalQuantile {
1513 coverage: 0.8,
1514 rank: 1,
1515 radii: vec![finite(2.0), finite(2.0)],
1516 },
1517 SplitConformalQuantile {
1518 coverage: 0.9,
1519 rank: 2,
1520 radii: vec![finite(3.0), finite(3.0)],
1521 },
1522 ];
1523 let marginal_intervals = apply_split_absolute_residual(
1524 &points,
1525 &marginal_quantiles,
1526 ConformalMultiTargetPolicy::Marginal,
1527 )
1528 .unwrap();
1529 assert_eq!(
1530 marginal_intervals[0].cells,
1531 vec![
1532 vec![finite_cell(8.0, 12.0), finite_cell(18.0, 22.0)],
1533 vec![finite_cell(9.0, 13.0), finite_cell(19.0, 23.0)],
1534 ]
1535 );
1536 let marginal_80 = regression_conformal_metrics(
1537 &truth,
1538 &marginal_intervals[0],
1539 ConformalMultiTargetPolicy::Marginal,
1540 )
1541 .unwrap();
1542 assert_close(marginal_80[0].empirical_coverage, 0.5);
1543 assert_close(marginal_80[0].coverage_gap, -0.300_000_000_000_000_04);
1544 assert_close(marginal_80[0].mean_width.unwrap(), 4.0);
1545 assert_close(marginal_80[0].median_width.unwrap(), 4.0);
1546 assert_close(marginal_80[0].interval_score.unwrap(), 9.0);
1547 assert_close(marginal_80[1].empirical_coverage, 1.0);
1548 assert_close(marginal_80[1].interval_score.unwrap(), 4.0);
1549 let marginal_90 = regression_conformal_metrics(
1550 &truth,
1551 &marginal_intervals[1],
1552 ConformalMultiTargetPolicy::Marginal,
1553 )
1554 .unwrap();
1555 for metric in marginal_90 {
1556 assert_close(metric.empirical_coverage, 1.0);
1557 assert_close(metric.mean_width.unwrap(), 6.0);
1558 assert_close(metric.interval_score.unwrap(), 6.0);
1559 }
1560
1561 let joint_quantiles = vec![
1562 SplitConformalQuantile {
1563 coverage: 0.8,
1564 rank: 1,
1565 radii: vec![finite(3.0)],
1566 },
1567 SplitConformalQuantile {
1568 coverage: 0.9,
1569 rank: 2,
1570 radii: vec![finite(4.0)],
1571 },
1572 ];
1573 let joint_intervals = apply_split_absolute_residual(
1574 &points,
1575 &joint_quantiles,
1576 ConformalMultiTargetPolicy::JointMax,
1577 )
1578 .unwrap();
1579 for (interval, expected_width) in joint_intervals.iter().zip([6.0, 8.0]) {
1580 let metric = regression_conformal_metrics(
1581 &truth,
1582 interval,
1583 ConformalMultiTargetPolicy::JointMax,
1584 )
1585 .unwrap();
1586 assert_close(metric[0].empirical_coverage, 1.0);
1587 assert_close(metric[0].mean_width.unwrap(), expected_width);
1588 assert_close(metric[0].median_width.unwrap(), expected_width);
1589 assert_close(metric[0].interval_score.unwrap(), expected_width);
1590 }
1591 }
1592
1593 #[test]
1594 fn unbounded_metrics_keep_coverage_and_tag_measurements_unavailable() {
1595 let truth = vec![vec![1.0, 10.0], vec![3.0, 20.0]];
1596 let interval = RegressionConformalInterval {
1597 coverage: 0.9,
1598 cells: vec![
1599 vec![RegressionIntervalCell::Unbounded, finite_cell(9.0, 11.0)],
1600 vec![RegressionIntervalCell::Unbounded, finite_cell(19.0, 21.0)],
1601 ],
1602 };
1603 let marginal =
1604 regression_conformal_metrics(&truth, &interval, ConformalMultiTargetPolicy::Marginal)
1605 .unwrap();
1606 assert_eq!(
1607 marginal[0].measurement_status,
1608 ConformalMeasurementStatus::Unbounded
1609 );
1610 assert_eq!(marginal[0].empirical_coverage, 1.0);
1611 assert_eq!(marginal[0].mean_width, None);
1612 assert_eq!(
1613 marginal[1].measurement_status,
1614 ConformalMeasurementStatus::Finite
1615 );
1616
1617 let joint =
1618 regression_conformal_metrics(&truth, &interval, ConformalMultiTargetPolicy::JointMax)
1619 .unwrap();
1620 assert_eq!(
1621 joint[0].measurement_status,
1622 ConformalMeasurementStatus::Unbounded
1623 );
1624 assert_eq!(joint[0].empirical_coverage, 1.0);
1625 assert_eq!(joint[0].interval_score, None);
1626 }
1627
1628 #[test]
1629 fn finite_metric_means_do_not_overflow_when_the_mean_is_representable() {
1630 let interval = RegressionConformalInterval {
1631 coverage: 0.5,
1632 cells: vec![
1633 vec![finite_cell(0.0, f64::MAX)],
1634 vec![finite_cell(0.0, f64::MAX)],
1635 ],
1636 };
1637 let metrics = regression_conformal_metrics(
1638 &[vec![0.0], vec![f64::MAX]],
1639 &interval,
1640 ConformalMultiTargetPolicy::Marginal,
1641 )
1642 .unwrap();
1643 assert_eq!(metrics[0].mean_width, Some(f64::MAX));
1644 assert_eq!(metrics[0].median_width, Some(f64::MAX));
1645 assert_eq!(metrics[0].interval_score, Some(f64::MAX));
1646 }
1647
1648 #[test]
1649 fn finite_metric_mean_keeps_w0_sequential_sum_rounding() {
1650 let values = [1.0, 1.0, 1.0e16];
1651 let mut sequential_sum = 0.0;
1652 for value in values {
1653 sequential_sum += value;
1654 }
1655 assert_eq!(sequential_sum / 3.0, 3_333_333_333_333_334.0);
1656 assert_eq!(
1657 checked_mean(&values, "rounding parity").unwrap(),
1658 sequential_sum / 3.0
1659 );
1660 }
1661
1662 #[test]
1663 fn invalid_coverages_and_residual_matrices_are_rejected() {
1664 for coverages in [
1665 vec![],
1666 vec![0.0],
1667 vec![1.0],
1668 vec![f64::NAN],
1669 vec![0.9, 0.8],
1670 vec![0.9, 0.9],
1671 ] {
1672 assert!(validate_conformal_coverages(&coverages).is_err());
1673 }
1674 let invalid = [
1675 vec![],
1676 vec![vec![]],
1677 vec![vec![1.0], vec![1.0, 2.0]],
1678 vec![vec![f64::INFINITY]],
1679 vec![vec![-1.0]],
1680 ];
1681 for residuals in invalid {
1682 assert!(split_absolute_residual_quantiles(
1683 &residuals,
1684 &[0.5],
1685 ConformalMultiTargetPolicy::Marginal,
1686 ConformalSmallSamplePolicy::Error,
1687 )
1688 .is_err());
1689 }
1690 }
1691
1692 #[test]
1693 fn application_rejects_bad_shape_status_order_and_overflow() {
1694 let points = vec![vec![1.0, 2.0]];
1695 let bad_shape = vec![SplitConformalQuantile {
1696 coverage: 0.8,
1697 rank: 2,
1698 radii: vec![finite(1.0)],
1699 }];
1700 assert!(apply_split_absolute_residual(
1701 &points,
1702 &bad_shape,
1703 ConformalMultiTargetPolicy::Marginal
1704 )
1705 .is_err());
1706
1707 let mixed = vec![SplitConformalQuantile {
1708 coverage: 0.8,
1709 rank: 2,
1710 radii: vec![finite(1.0), ConformalRadius::Unbounded],
1711 }];
1712 assert!(matches!(
1713 apply_split_absolute_residual(&points, &mixed, ConformalMultiTargetPolicy::Marginal),
1714 Err(ConformalError::MixedRadiusStatus { .. })
1715 ));
1716
1717 let non_nested = vec![
1718 SplitConformalQuantile {
1719 coverage: 0.8,
1720 rank: 2,
1721 radii: vec![finite(2.0), finite(2.0)],
1722 },
1723 SplitConformalQuantile {
1724 coverage: 0.9,
1725 rank: 3,
1726 radii: vec![finite(1.0), finite(2.0)],
1727 },
1728 ];
1729 assert!(matches!(
1730 apply_split_absolute_residual(
1731 &points,
1732 &non_nested,
1733 ConformalMultiTargetPolicy::Marginal
1734 ),
1735 Err(ConformalError::NonNestedRadius { .. })
1736 ));
1737
1738 let overflow = vec![SplitConformalQuantile {
1739 coverage: 0.8,
1740 rank: 2,
1741 radii: vec![finite(f64::MAX)],
1742 }];
1743 assert!(matches!(
1744 apply_split_absolute_residual(
1745 &[vec![f64::MAX]],
1746 &overflow,
1747 ConformalMultiTargetPolicy::Marginal
1748 ),
1749 Err(ConformalError::ArithmeticOverflow { .. })
1750 ));
1751 }
1752
1753 #[test]
1754 fn metrics_reject_invalid_shapes_bounds_and_nonfinite_truth() {
1755 let valid_interval = RegressionConformalInterval {
1756 coverage: 0.8,
1757 cells: vec![vec![finite_cell(0.0, 2.0)]],
1758 };
1759 assert!(regression_conformal_metrics(
1760 &[vec![f64::NAN]],
1761 &valid_interval,
1762 ConformalMultiTargetPolicy::Marginal
1763 )
1764 .is_err());
1765 assert!(regression_conformal_metrics(
1766 &[vec![1.0], vec![2.0]],
1767 &valid_interval,
1768 ConformalMultiTargetPolicy::Marginal
1769 )
1770 .is_err());
1771 let bad_bounds = RegressionConformalInterval {
1772 coverage: 0.8,
1773 cells: vec![vec![finite_cell(2.0, 1.0)]],
1774 };
1775 assert!(matches!(
1776 regression_conformal_metrics(
1777 &[vec![1.0]],
1778 &bad_bounds,
1779 ConformalMultiTargetPolicy::Marginal
1780 ),
1781 Err(ConformalError::InvalidIntervalCell { .. })
1782 ));
1783 }
1784}