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