1use std::collections::{BTreeMap, BTreeSet};
2
3use serde::{Deserialize, Serialize};
4
5use crate::aggregation::{
6 reduce_predictions_across_folds, AggregatedPredictionBlock, PredictionUnitId,
7};
8use crate::error::{DagMlError, Result};
9use crate::fold::FoldPartitionMode;
10use crate::ids::{FoldId, NodeId, SampleId, VariantId};
11use crate::metric_provider::{
12 builtin_metric_reference, builtin_metric_registry, MetricEvaluationScope, MetricEvaluationTask,
13 MetricUnitId,
14};
15use crate::oof::{validate_producer_oof_coverage, PredictionBlock, PredictionPartition};
16use crate::policy::PredictionLevel;
17use crate::selection::{CandidateScore, MetricObjective};
18use crate::{LearningTaskKind, PredictionKind};
19
20#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Serialize, Deserialize)]
21#[serde(rename_all = "snake_case")]
22pub enum RegressionMetricKind {
23 Mse,
24 Rmse,
25 Mae,
26 R2,
27 Accuracy,
31 BalancedAccuracy,
39}
40
41impl RegressionMetricKind {
42 pub fn from_name(name: &str) -> Option<Self> {
43 match name {
44 "mse" => Some(Self::Mse),
45 "rmse" => Some(Self::Rmse),
46 "mae" => Some(Self::Mae),
47 "r2" => Some(Self::R2),
48 "accuracy" => Some(Self::Accuracy),
49 "balanced_accuracy" => Some(Self::BalancedAccuracy),
50 _ => None,
51 }
52 }
53
54 pub fn name(self) -> &'static str {
55 match self {
56 Self::Mse => "mse",
57 Self::Rmse => "rmse",
58 Self::Mae => "mae",
59 Self::R2 => "r2",
60 Self::Accuracy => "accuracy",
61 Self::BalancedAccuracy => "balanced_accuracy",
62 }
63 }
64
65 pub fn objective(self) -> MetricObjective {
66 match self {
67 Self::Mse | Self::Rmse | Self::Mae => MetricObjective::Minimize,
68 Self::R2 | Self::Accuracy | Self::BalancedAccuracy => MetricObjective::Maximize,
69 }
70 }
71
72 pub fn resolve_for_prediction_kind(
76 name: &str,
77 objective: MetricObjective,
78 prediction_kind: crate::training::PredictionKind,
79 ) -> Result<Self> {
80 let metric = Self::from_name(name).ok_or_else(|| {
81 DagMlError::CampaignValidation(format!("unsupported native selection metric `{name}`"))
82 })?;
83 let kind_compatible = match prediction_kind {
84 crate::training::PredictionKind::RegressionPoint => {
85 matches!(metric, Self::Mse | Self::Rmse | Self::Mae | Self::R2)
86 }
87 crate::training::PredictionKind::ClassLabel => {
88 matches!(metric, Self::Accuracy | Self::BalancedAccuracy)
89 }
90 crate::training::PredictionKind::ClassProbability
91 | crate::training::PredictionKind::DecisionScore => false,
92 };
93 if objective != metric.objective() || !kind_compatible {
94 return Err(DagMlError::CampaignValidation(format!(
95 "selection metric `{name}` with objective {objective:?} is not supported for {prediction_kind:?} output"
96 )));
97 }
98 Ok(metric)
99 }
100}
101
102#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
103#[serde(deny_unknown_fields)]
104pub struct RegressionTargetBlock {
105 pub level: PredictionLevel,
106 pub unit_ids: Vec<PredictionUnitId>,
107 pub values: Vec<Vec<f64>>,
108 #[serde(default)]
109 pub target_names: Vec<String>,
110}
111
112impl RegressionTargetBlock {
113 pub fn validate_shape(&self) -> Result<usize> {
114 if self.unit_ids.len() != self.values.len() {
115 return Err(DagMlError::OofValidation(format!(
116 "target block has {} unit ids but {} target rows",
117 self.unit_ids.len(),
118 self.values.len()
119 )));
120 }
121 if self
122 .unit_ids
123 .iter()
124 .any(|unit_id| unit_id.level() != self.level)
125 {
126 return Err(DagMlError::OofValidation(format!(
127 "target block contains units outside level {:?}",
128 self.level
129 )));
130 }
131 let unique = self.unit_ids.iter().collect::<BTreeSet<_>>();
132 if unique.len() != self.unit_ids.len() {
133 return Err(DagMlError::OofValidation(
134 "target block contains duplicate unit ids".to_string(),
135 ));
136 }
137 let width = self.values.first().map_or(0, Vec::len);
138 if width == 0 {
139 return Err(DagMlError::OofValidation(
140 "target block has empty target rows".to_string(),
141 ));
142 }
143 if self.values.iter().any(|row| row.len() != width) {
144 return Err(DagMlError::OofValidation(
145 "target block has ragged target rows".to_string(),
146 ));
147 }
148 if self.values.iter().flatten().any(|value| !value.is_finite()) {
149 return Err(DagMlError::OofValidation(
150 "target block contains non-finite values".to_string(),
151 ));
152 }
153 if !self.target_names.is_empty() && self.target_names.len() != width {
154 return Err(DagMlError::OofValidation(format!(
155 "target block has {} target names for width {}",
156 self.target_names.len(),
157 width
158 )));
159 }
160 Ok(width)
161 }
162}
163
164pub fn reassemble_merge_targets(
188 producer_node: &NodeId,
189 merge_sample_ids: &[SampleId],
190 by_sample_target: &mut BTreeMap<SampleId, Vec<f64>>,
191 target_names: Vec<String>,
192) -> Result<Option<RegressionTargetBlock>> {
193 if by_sample_target.is_empty() {
194 return Ok(None);
195 }
196 let missing: Vec<String> = merge_sample_ids
197 .iter()
198 .filter(|sample_id| !by_sample_target.contains_key(*sample_id))
199 .map(ToString::to_string)
200 .collect();
201 if !missing.is_empty() {
202 return Err(DagMlError::OofValidation(format!(
203 "merge node `{producer_node}` has partial target coverage: {} of {} merged sample(s) lack a y_true row ({}) while other contributing branch(es) emitted targets — a merge that some branch scores must have COMPLETE target coverage across the merge universe, never a silent no-score",
204 missing.len(),
205 merge_sample_ids.len(),
206 missing.join(", ")
207 )));
208 }
209 let values: Vec<Vec<f64>> = merge_sample_ids
210 .iter()
211 .map(|sample_id| {
212 by_sample_target
213 .remove(sample_id)
214 .expect("target coverage was just verified complete")
215 })
216 .collect();
217 Ok(Some(RegressionTargetBlock {
218 level: PredictionLevel::Sample,
219 unit_ids: merge_sample_ids
220 .iter()
221 .cloned()
222 .map(PredictionUnitId::Sample)
223 .collect(),
224 values,
225 target_names,
226 }))
227}
228
229#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
230pub struct RegressionMetricReport {
231 #[serde(default)]
232 pub prediction_id: Option<String>,
233 pub producer_node: NodeId,
234 #[serde(default, skip_serializing_if = "Option::is_none")]
235 pub producer_port: Option<String>,
236 #[serde(default, skip_serializing_if = "Option::is_none")]
240 pub variant_id: Option<VariantId>,
241 #[serde(default, skip_serializing_if = "Option::is_none")]
249 pub variant_label: Option<String>,
250 pub partition: PredictionPartition,
251 pub fold_id: Option<FoldId>,
252 pub level: PredictionLevel,
253 pub row_count: usize,
254 pub target_width: usize,
255 #[serde(default)]
256 pub target_names: Vec<String>,
257 pub metrics: BTreeMap<String, f64>,
258}
259
260impl RegressionMetricReport {
261 pub fn validate(&self) -> Result<()> {
262 if self.row_count == 0 {
263 return Err(DagMlError::OofValidation(
264 "regression metric report has zero rows".to_string(),
265 ));
266 }
267 if self.target_width == 0 {
268 return Err(DagMlError::OofValidation(
269 "regression metric report has zero target width".to_string(),
270 ));
271 }
272 if !self.target_names.is_empty() && self.target_names.len() != self.target_width {
273 return Err(DagMlError::OofValidation(format!(
274 "regression metric report has {} target names for width {}",
275 self.target_names.len(),
276 self.target_width
277 )));
278 }
279 if self.metrics.is_empty() {
280 return Err(DagMlError::OofValidation(
281 "regression metric report has no metrics".to_string(),
282 ));
283 }
284 for (name, value) in &self.metrics {
285 if name.trim().is_empty() {
286 return Err(DagMlError::OofValidation(
287 "regression metric report contains an empty metric name".to_string(),
288 ));
289 }
290 if !value.is_finite() {
291 return Err(DagMlError::OofValidation(format!(
292 "regression metric `{name}` is not finite"
293 )));
294 }
295 }
296 Ok(())
297 }
298
299 pub fn into_candidate_score(self, candidate_id: impl Into<String>) -> Result<CandidateScore> {
300 self.validate()?;
301 let mut metadata = BTreeMap::from([
302 (
303 "producer_node".to_string(),
304 serde_json::json!(self.producer_node),
305 ),
306 ("partition".to_string(), serde_json::json!(self.partition)),
307 (
308 "metric_level".to_string(),
309 serde_json::json!(prediction_level_name(self.level)),
310 ),
311 ("row_count".to_string(), serde_json::json!(self.row_count)),
312 (
313 "target_width".to_string(),
314 serde_json::json!(self.target_width),
315 ),
316 ]);
317 if let Some(prediction_id) = self.prediction_id {
318 metadata.insert(
319 "prediction_id".to_string(),
320 serde_json::json!(prediction_id),
321 );
322 }
323 if let Some(producer_port) = self.producer_port {
324 metadata.insert(
325 "producer_port".to_string(),
326 serde_json::json!(producer_port),
327 );
328 }
329 if let Some(fold_id) = self.fold_id {
330 metadata.insert("fold_id".to_string(), serde_json::json!(fold_id));
331 }
332 if let Some(variant_id) = self.variant_id {
333 metadata.insert("variant_id".to_string(), serde_json::json!(variant_id));
334 }
335 if !self.target_names.is_empty() {
336 metadata.insert(
337 "target_names".to_string(),
338 serde_json::json!(self.target_names),
339 );
340 }
341 let score = CandidateScore {
342 candidate_id: candidate_id.into(),
343 metrics: self.metrics,
344 metadata,
345 };
346 score.validate()?;
347 Ok(score)
348 }
349}
350
351pub fn regression_report_to_candidate_score(
352 candidate_id: impl Into<String>,
353 report: RegressionMetricReport,
354) -> Result<CandidateScore> {
355 report.into_candidate_score(candidate_id)
356}
357
358pub fn score_regression_prediction_block(
359 predictions: &PredictionBlock,
360 targets: &RegressionTargetBlock,
361 metrics: &[RegressionMetricKind],
362) -> Result<RegressionMetricReport> {
363 let width = validate_sample_prediction_block(predictions)?;
364 let prediction_units = predictions
365 .sample_ids
366 .iter()
367 .cloned()
368 .map(PredictionUnitId::Sample)
369 .collect::<Vec<_>>();
370 score_regression_rows(
371 PredictionRows {
372 level: PredictionLevel::Sample,
373 unit_ids: &prediction_units,
374 values: &predictions.values,
375 target_names: &predictions.target_names,
376 width,
377 origin: PredictionReportOrigin {
378 prediction_id: predictions.prediction_id.clone(),
379 producer_node: predictions.producer_node.clone(),
380 producer_port: predictions.producer_port.clone(),
381 partition: predictions.partition.clone(),
382 fold_id: predictions.fold_id.clone(),
383 },
384 },
385 targets,
386 metrics,
387 )
388}
389
390pub fn score_regression_aggregated_block(
391 predictions: &AggregatedPredictionBlock,
392 targets: &RegressionTargetBlock,
393 metrics: &[RegressionMetricKind],
394) -> Result<RegressionMetricReport> {
395 let width = predictions.validate_shape()?;
396 score_regression_rows(
397 PredictionRows {
398 level: predictions.level,
399 unit_ids: &predictions.unit_ids,
400 values: &predictions.values,
401 target_names: &predictions.target_names,
402 width,
403 origin: PredictionReportOrigin {
404 prediction_id: predictions.prediction_id.clone(),
405 producer_node: predictions.producer_node.clone(),
406 producer_port: predictions.producer_port.clone(),
407 partition: predictions.partition.clone(),
408 fold_id: predictions.fold_id.clone(),
409 },
410 },
411 targets,
412 metrics,
413 )
414}
415
416pub const SCORE_SET_SCHEMA_VERSION: u32 = 2;
418pub const LEGACY_SCORE_SET_SCHEMA_VERSION: u32 = 1;
419pub const MIN_READABLE_SCORE_SET_SCHEMA_VERSION: u32 = 1;
420
421fn default_score_set_schema_version() -> u32 {
422 LEGACY_SCORE_SET_SCHEMA_VERSION
423}
424
425type ScoreReportKey = (
436 NodeId,
437 Option<String>,
438 Option<VariantId>,
439 PredictionPartition,
440 Option<FoldId>,
441 PredictionLevel,
442);
443
444#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
445pub struct ScoreSet {
446 #[serde(default = "default_score_set_schema_version")]
447 pub schema_version: u32,
448 pub plan_id: String,
449 #[serde(default, skip_serializing_if = "Option::is_none")]
451 pub selection_metric: Option<String>,
452 pub reports: Vec<RegressionMetricReport>,
453}
454
455impl ScoreSet {
456 pub fn validate(&self) -> Result<()> {
458 if self.schema_version < MIN_READABLE_SCORE_SET_SCHEMA_VERSION
459 || self.schema_version > SCORE_SET_SCHEMA_VERSION
460 {
461 return Err(DagMlError::OofValidation(format!(
462 "score set schema version {} is unsupported (current {SCORE_SET_SCHEMA_VERSION})",
463 self.schema_version
464 )));
465 }
466 if self.plan_id.trim().is_empty() {
467 return Err(DagMlError::OofValidation(
468 "score set has an empty plan_id".to_string(),
469 ));
470 }
471 let mut seen: BTreeSet<ScoreReportKey> = BTreeSet::new();
472 for report in &self.reports {
473 report.validate()?;
474 match (self.schema_version, report.producer_port.as_deref()) {
475 (LEGACY_SCORE_SET_SCHEMA_VERSION, Some(_)) => {
476 return Err(DagMlError::OofValidation(
477 "score set V1 reports must not carry producer_port".to_string(),
478 ));
479 }
480 (SCORE_SET_SCHEMA_VERSION, Some(port)) if port.trim().is_empty() => {
481 return Err(DagMlError::OofValidation(
482 "score set V2 report has an empty producer_port".to_string(),
483 ));
484 }
485 (SCORE_SET_SCHEMA_VERSION, None) => {
486 return Err(DagMlError::OofValidation(
487 "score set V2 requires producer_port on every report".to_string(),
488 ));
489 }
490 _ => {}
491 }
492 let key = (
493 report.producer_node.clone(),
494 report.producer_port.clone(),
495 report.variant_id.clone(),
496 report.partition.clone(),
497 report.fold_id.clone(),
498 report.level,
499 );
500 if !seen.insert(key) {
501 return Err(DagMlError::OofValidation(format!(
502 "score set has a duplicate report for node `{}` port {:?} partition {:?} fold {:?} level {:?}",
503 report.producer_node,
504 report.producer_port,
505 report.partition,
506 report.fold_id,
507 report.level
508 )));
509 }
510 }
511 Ok(())
512 }
513}
514
515#[derive(Clone, Debug)]
516struct PredictionReportOrigin {
517 prediction_id: Option<String>,
518 producer_node: NodeId,
519 producer_port: Option<String>,
520 partition: PredictionPartition,
521 fold_id: Option<FoldId>,
522}
523
524#[derive(Clone, Debug)]
525struct PredictionRows<'a> {
526 level: PredictionLevel,
527 unit_ids: &'a [PredictionUnitId],
528 values: &'a [Vec<f64>],
529 target_names: &'a [String],
530 width: usize,
531 origin: PredictionReportOrigin,
532}
533
534fn score_regression_rows(
535 predictions: PredictionRows<'_>,
536 targets: &RegressionTargetBlock,
537 metrics: &[RegressionMetricKind],
538) -> Result<RegressionMetricReport> {
539 if metrics.is_empty() {
540 return Err(DagMlError::OofValidation(
541 "no regression metrics requested".to_string(),
542 ));
543 }
544 let mut requested_metrics = BTreeSet::new();
545 for metric in metrics {
546 if !requested_metrics.insert(*metric) {
547 return Err(DagMlError::OofValidation(format!(
548 "duplicate regression metric `{}` requested",
549 metric.name()
550 )));
551 }
552 }
553
554 let target_width = targets.validate_shape()?;
555 if predictions.width != target_width {
556 return Err(DagMlError::OofValidation(format!(
557 "prediction width {} does not match target width {target_width}",
558 predictions.width
559 )));
560 }
561 if predictions.level != targets.level {
562 return Err(DagMlError::OofValidation(format!(
563 "prediction level {:?} does not match target level {:?}",
564 predictions.level, targets.level
565 )));
566 }
567 if !predictions.target_names.is_empty()
568 && !targets.target_names.is_empty()
569 && predictions.target_names != targets.target_names
570 {
571 return Err(DagMlError::OofValidation(
572 "prediction target names do not match target block names".to_string(),
573 ));
574 }
575
576 let target_by_unit = targets
577 .unit_ids
578 .iter()
579 .zip(targets.values.iter().map(Vec::as_slice))
580 .collect::<BTreeMap<_, _>>();
581 let mut aligned_predictions = Vec::with_capacity(predictions.unit_ids.len());
582 let mut aligned_targets = Vec::with_capacity(predictions.unit_ids.len());
583 for (unit_id, prediction_row) in predictions.unit_ids.iter().zip(predictions.values.iter()) {
584 let target_row = target_by_unit.get(unit_id).ok_or_else(|| {
585 DagMlError::OofValidation(format!(
586 "prediction unit `{unit_id}` is missing from target block"
587 ))
588 })?;
589 aligned_predictions.push(prediction_row.as_slice());
590 aligned_targets.push(*target_row);
591 }
592 if aligned_predictions.len() != target_by_unit.len() {
593 return Err(DagMlError::OofValidation(
594 "target block contains units not present in predictions".to_string(),
595 ));
596 }
597
598 let target_names = if !predictions.target_names.is_empty() {
599 predictions.target_names.to_vec()
600 } else {
601 targets.target_names.clone()
602 };
603 let metric_suffixes = target_metric_names(predictions.width, &target_names);
604 let provider_output_ids = (0..predictions.width)
605 .map(|index| format!("output:{index}"))
606 .collect::<Vec<_>>();
607 let metric_units = predictions
608 .unit_ids
609 .iter()
610 .map(MetricUnitId::from)
611 .collect::<Vec<_>>();
612 let prediction_values = aligned_predictions
613 .iter()
614 .map(|row| row.to_vec())
615 .collect::<Vec<_>>();
616 let target_values = aligned_targets
617 .iter()
618 .map(|row| row.to_vec())
619 .collect::<Vec<_>>();
620 let scope = MetricEvaluationScope {
621 producer_node: predictions.origin.producer_node.clone(),
622 producer_port: predictions.origin.producer_port.clone(),
623 prediction_id: predictions.origin.prediction_id.clone(),
624 variant_id: None,
625 partition: predictions.origin.partition.clone(),
626 fold_id: predictions.origin.fold_id.clone(),
627 level: predictions.level,
628 };
629 let registry = builtin_metric_registry()?;
630 let mut values = BTreeMap::new();
631 for metric in metrics {
632 let (task_kind, prediction_kind) = match metric {
633 RegressionMetricKind::Mse
634 | RegressionMetricKind::Rmse
635 | RegressionMetricKind::Mae
636 | RegressionMetricKind::R2 => (
637 LearningTaskKind::Regression,
638 PredictionKind::RegressionPoint,
639 ),
640 RegressionMetricKind::Accuracy | RegressionMetricKind::BalancedAccuracy => (
641 LearningTaskKind::MulticlassClassification,
642 PredictionKind::ClassLabel,
643 ),
644 };
645 let task = MetricEvaluationTask::new(
646 format!(
647 "metric:{}:{}",
648 predictions.origin.producer_node,
649 metric.name()
650 ),
651 builtin_metric_reference(*metric)?,
652 task_kind,
653 prediction_kind,
654 scope.clone(),
655 metric_units.clone(),
656 prediction_values.clone(),
657 target_values.clone(),
658 provider_output_ids.clone(),
659 None,
660 None,
661 None,
662 )?;
663 let evaluation = registry.evaluate(&task)?;
664 values.insert(metric.name().to_string(), evaluation.aggregate);
665 for (component, suffix) in evaluation.result.values.into_iter().zip(&metric_suffixes) {
666 values.insert(format!("{}:{suffix}", metric.name()), component.value);
667 }
668 }
669
670 let report = RegressionMetricReport {
671 prediction_id: predictions.origin.prediction_id,
672 producer_node: predictions.origin.producer_node,
673 producer_port: predictions.origin.producer_port,
674 variant_id: None,
675 variant_label: None,
676 partition: predictions.origin.partition,
677 fold_id: predictions.origin.fold_id,
678 level: predictions.level,
679 row_count: predictions.unit_ids.len(),
680 target_width: predictions.width,
681 target_names,
682 metrics: values,
683 };
684 report.validate()?;
685 Ok(report)
686}
687
688fn validate_sample_prediction_block(block: &PredictionBlock) -> Result<usize> {
689 block.validate_content()
690}
691
692pub(crate) fn compute_metric_per_target(
693 metric: RegressionMetricKind,
694 width: usize,
695 predictions: &[&[f64]],
696 targets: &[&[f64]],
697) -> Vec<f64> {
698 (0..width)
699 .map(|target_idx| match metric {
700 RegressionMetricKind::Mse => {
701 predictions
702 .iter()
703 .zip(targets.iter())
704 .map(|(prediction, target)| {
705 let error = prediction[target_idx] - target[target_idx];
706 error * error
707 })
708 .sum::<f64>()
709 / predictions.len() as f64
710 }
711 RegressionMetricKind::Rmse => (predictions
712 .iter()
713 .zip(targets.iter())
714 .map(|(prediction, target)| {
715 let error = prediction[target_idx] - target[target_idx];
716 error * error
717 })
718 .sum::<f64>()
719 / predictions.len() as f64)
720 .sqrt(),
721 RegressionMetricKind::Mae => {
722 predictions
723 .iter()
724 .zip(targets.iter())
725 .map(|(prediction, target)| (prediction[target_idx] - target[target_idx]).abs())
726 .sum::<f64>()
727 / predictions.len() as f64
728 }
729 RegressionMetricKind::R2 => r2_for_target(target_idx, predictions, targets),
730 RegressionMetricKind::Accuracy => {
731 predictions
732 .iter()
733 .zip(targets.iter())
734 .filter(|(prediction, target)| {
735 (prediction[target_idx] - target[target_idx]).abs() < 0.5
736 })
737 .count() as f64
738 / predictions.len() as f64
739 }
740 RegressionMetricKind::BalancedAccuracy => {
741 balanced_accuracy_for_target(target_idx, predictions, targets)
742 }
743 })
744 .collect()
745}
746
747fn balanced_accuracy_for_target(
754 target_idx: usize,
755 predictions: &[&[f64]],
756 targets: &[&[f64]],
757) -> f64 {
758 let mut per_class: BTreeMap<i64, (usize, usize)> = BTreeMap::new();
760 for (prediction, target) in predictions.iter().zip(targets.iter()) {
761 let true_value = target[target_idx];
762 let class = true_value.round() as i64;
763 let entry = per_class.entry(class).or_insert((0, 0));
764 entry.1 += 1;
765 if (prediction[target_idx] - true_value).abs() < 0.5 {
766 entry.0 += 1;
767 }
768 }
769 if per_class.is_empty() {
770 return 0.0;
771 }
772 let recall_sum: f64 = per_class
773 .values()
774 .map(|(correct, count)| *correct as f64 / *count as f64)
775 .sum();
776 recall_sum / per_class.len() as f64
777}
778
779fn r2_for_target(target_idx: usize, predictions: &[&[f64]], targets: &[&[f64]]) -> f64 {
780 let mean = targets.iter().map(|row| row[target_idx]).sum::<f64>() / targets.len() as f64;
781 let ss_res = predictions
782 .iter()
783 .zip(targets.iter())
784 .map(|(prediction, target)| {
785 let error = prediction[target_idx] - target[target_idx];
786 error * error
787 })
788 .sum::<f64>();
789 let ss_tot = targets
790 .iter()
791 .map(|target| {
792 let centered = target[target_idx] - mean;
793 centered * centered
794 })
795 .sum::<f64>();
796 if ss_tot == 0.0 {
797 if ss_res == 0.0 {
798 1.0
799 } else {
800 0.0
801 }
802 } else {
803 1.0 - ss_res / ss_tot
804 }
805}
806
807fn target_metric_names(width: usize, target_names: &[String]) -> Vec<String> {
808 if target_names.is_empty() {
809 (0..width).map(|idx| format!("target_{idx}")).collect()
810 } else {
811 target_names.to_vec()
812 }
813}
814
815fn prediction_level_name(level: PredictionLevel) -> &'static str {
816 match level {
817 PredictionLevel::Observation => "observation",
818 PredictionLevel::Sample => "sample",
819 PredictionLevel::Target => "target",
820 PredictionLevel::Group => "group",
821 }
822}
823
824#[derive(Clone, Debug, PartialEq)]
827pub struct RegressionTargetRecord {
828 pub producer_node: NodeId,
829 pub producer_port: Option<String>,
830 pub variant_id: Option<VariantId>,
833 pub partition: PredictionPartition,
834 pub fold_id: Option<FoldId>,
835 pub block: RegressionTargetBlock,
836}
837
838fn combine_validation_targets(
847 producer: &NodeId,
848 producer_port: &Option<String>,
849 records: &[RegressionTargetRecord],
850) -> Result<RegressionTargetBlock> {
851 let mut seen: BTreeMap<PredictionUnitId, Vec<f64>> = BTreeMap::new();
852 let mut unit_ids = Vec::new();
853 let mut values = Vec::new();
854 let mut target_names = Vec::new();
855 for record in records {
856 if &record.producer_node != producer
857 || &record.producer_port != producer_port
858 || record.partition != PredictionPartition::Validation
859 {
860 continue;
861 }
862 if target_names.is_empty() {
863 target_names = record.block.target_names.clone();
864 }
865 for (unit_id, row) in record.block.unit_ids.iter().zip(&record.block.values) {
866 match seen.get(unit_id) {
867 None => {
868 seen.insert(unit_id.clone(), row.clone());
869 unit_ids.push(unit_id.clone());
870 values.push(row.clone());
871 }
872 Some(existing) if existing != row => {
873 return Err(DagMlError::OofValidation(format!(
874 "producer `{producer}` has conflicting ground truth for unit `{unit_id:?}` across validation records — the y_true reference is mixed (e.g. several variants in one context); refusing to score against a corrupted reference"
875 )));
876 }
877 Some(_) => {}
878 }
879 }
880 }
881 Ok(RegressionTargetBlock {
882 level: PredictionLevel::Sample,
883 unit_ids,
884 values,
885 target_names,
886 })
887}
888
889#[derive(Clone, Debug, PartialEq)]
898pub struct OofAverageBlock {
899 pub predictions: AggregatedPredictionBlock,
900 pub y_true: RegressionTargetBlock,
901}
902
903#[derive(Clone, Debug, Default, PartialEq)]
908pub struct CrossFoldValidation {
909 pub reports: Vec<RegressionMetricReport>,
910 pub oof_averages: Vec<OofAverageBlock>,
911}
912
913pub fn cross_fold_validation_reports(
927 prediction_blocks: &[PredictionBlock],
928 target_records: &[RegressionTargetRecord],
929 metrics: &[RegressionMetricKind],
930 partition_mode: FoldPartitionMode,
931) -> Result<CrossFoldValidation> {
932 let mut producers: Vec<(NodeId, Option<String>)> = Vec::new();
933 let mut by_producer: BTreeMap<(NodeId, Option<String>), Vec<PredictionBlock>> = BTreeMap::new();
934 for block in prediction_blocks {
935 if block.partition != PredictionPartition::Validation {
936 continue;
937 }
938 let key = (block.producer_node.clone(), block.producer_port.clone());
939 if !by_producer.contains_key(&key) {
940 producers.push(key.clone());
941 }
942 by_producer.entry(key).or_default().push(block.clone());
943 }
944 let mut reports = Vec::new();
945 let mut oof_averages = Vec::new();
946 for (producer, producer_port) in &producers {
947 let blocks = &by_producer[&(producer.clone(), producer_port.clone())];
948 if blocks.len() < 2 {
949 continue;
950 }
951 let block_refs = blocks.iter().collect::<Vec<_>>();
961 validate_producer_oof_coverage(producer, &block_refs, partition_mode, None)?;
962 let targets = combine_validation_targets(producer, producer_port, target_records)?;
963 if targets.unit_ids.is_empty() {
964 continue;
966 }
967 let average = reduce_predictions_across_folds(blocks, None, "avg")?;
968 reports.push(score_regression_prediction_block(
972 &average, &targets, metrics,
973 )?);
974 oof_averages.push(oof_average_block(&average, &targets));
975 }
976 Ok(CrossFoldValidation {
977 reports,
978 oof_averages,
979 })
980}
981
982fn oof_average_block(
990 average: &PredictionBlock,
991 targets: &RegressionTargetBlock,
992) -> OofAverageBlock {
993 let unit_ids: Vec<PredictionUnitId> = average
994 .sample_ids
995 .iter()
996 .cloned()
997 .map(PredictionUnitId::Sample)
998 .collect();
999 let predictions = AggregatedPredictionBlock {
1000 prediction_id: None,
1001 producer_node: average.producer_node.clone(),
1002 producer_port: average.producer_port.clone(),
1003 partition: average.partition.clone(),
1004 fold_id: average.fold_id.clone(),
1005 level: PredictionLevel::Sample,
1006 unit_ids: unit_ids.clone(),
1007 values: average.values.clone(),
1008 target_names: average.target_names.clone(),
1009 };
1010 let target_by_unit: BTreeMap<&PredictionUnitId, &Vec<f64>> =
1011 targets.unit_ids.iter().zip(&targets.values).collect();
1012 let y_true = RegressionTargetBlock {
1013 level: PredictionLevel::Sample,
1014 unit_ids: unit_ids.clone(),
1015 values: unit_ids
1016 .iter()
1017 .map(|unit_id| target_by_unit[unit_id].clone())
1018 .collect(),
1019 target_names: targets.target_names.clone(),
1020 };
1021 OofAverageBlock {
1022 predictions,
1023 y_true,
1024 }
1025}
1026
1027#[cfg(test)]
1028mod tests {
1029 use super::*;
1030 use crate::ids::{FoldId, GroupId, NodeId, SampleId, TargetId};
1031 use crate::oof::PredictionPartition;
1032
1033 fn sid(value: &str) -> SampleId {
1034 SampleId::new(value).unwrap()
1035 }
1036
1037 fn sample_unit(value: &str) -> PredictionUnitId {
1038 PredictionUnitId::Sample(sid(value))
1039 }
1040
1041 fn target_unit(value: &str) -> PredictionUnitId {
1042 PredictionUnitId::Target(TargetId::new(value).unwrap())
1043 }
1044
1045 fn group_unit(value: &str) -> PredictionUnitId {
1046 PredictionUnitId::Group(GroupId::new(value).unwrap())
1047 }
1048
1049 fn assert_close(left: f64, right: f64) {
1050 assert!((left - right).abs() < 1e-12, "expected {right}, got {left}");
1051 }
1052
1053 #[test]
1054 fn metric_objectives_match_selection_direction() {
1055 assert_eq!(
1056 RegressionMetricKind::Rmse.objective(),
1057 MetricObjective::Minimize
1058 );
1059 assert_eq!(
1060 RegressionMetricKind::Mae.objective(),
1061 MetricObjective::Minimize
1062 );
1063 assert_eq!(
1064 RegressionMetricKind::Mse.objective(),
1065 MetricObjective::Minimize
1066 );
1067 assert_eq!(
1068 RegressionMetricKind::R2.objective(),
1069 MetricObjective::Maximize
1070 );
1071 }
1072
1073 #[test]
1074 fn reassemble_merge_targets_empty_map_is_unscored_none() {
1075 let producer = NodeId::new("merge:m").unwrap();
1077 let mut by_sample: BTreeMap<SampleId, Vec<f64>> = BTreeMap::new();
1078 let block = reassemble_merge_targets(
1079 &producer,
1080 &[sid("s1"), sid("s2")],
1081 &mut by_sample,
1082 vec!["y".to_string()],
1083 )
1084 .unwrap();
1085 assert!(
1086 block.is_none(),
1087 "empty targets -> unscored None, not an error"
1088 );
1089 }
1090
1091 #[test]
1092 fn reassemble_merge_targets_complete_coverage_emits_ordered_block() {
1093 let producer = NodeId::new("merge:m").unwrap();
1094 let mut by_sample: BTreeMap<SampleId, Vec<f64>> = BTreeMap::new();
1095 by_sample.insert(sid("s2"), vec![20.0]);
1096 by_sample.insert(sid("s1"), vec![10.0]);
1097 let block = reassemble_merge_targets(
1098 &producer,
1099 &[sid("s1"), sid("s2")],
1100 &mut by_sample,
1101 vec!["y".to_string()],
1102 )
1103 .unwrap()
1104 .expect("complete coverage -> a target block");
1105 assert_eq!(
1107 block.unit_ids,
1108 vec![sample_unit("s1"), sample_unit("s2")],
1109 "targets follow the merge sample order"
1110 );
1111 assert_eq!(block.values, vec![vec![10.0], vec![20.0]]);
1112 assert_eq!(block.level, PredictionLevel::Sample);
1113 block.validate_shape().unwrap();
1114 }
1115
1116 #[test]
1117 fn reassemble_merge_targets_partial_coverage_is_validation_error() {
1118 let producer = NodeId::new("merge:m").unwrap();
1122 let mut by_sample: BTreeMap<SampleId, Vec<f64>> = BTreeMap::new();
1123 by_sample.insert(sid("s1"), vec![10.0]);
1124 let err = reassemble_merge_targets(
1125 &producer,
1126 &[sid("s1"), sid("s2")],
1127 &mut by_sample,
1128 vec!["y".to_string()],
1129 )
1130 .unwrap_err();
1131 let msg = err.to_string();
1132 assert!(
1133 msg.contains("partial target coverage") && msg.contains("s2"),
1134 "partial coverage names the missing sample: {msg}"
1135 );
1136 }
1137
1138 #[test]
1139 fn scores_sample_predictions_and_exports_candidate_metrics() {
1140 let predictions = PredictionBlock {
1141 prediction_id: Some("pred:sample".to_string()),
1142 producer_node: NodeId::new("model:pls").unwrap(),
1143 producer_port: None,
1144 partition: PredictionPartition::Validation,
1145 fold_id: None,
1146 sample_ids: vec![sid("sample:1"), sid("sample:2")],
1147 values: vec![vec![2.0], vec![4.0]],
1148 target_names: vec!["y".to_string()],
1149 };
1150 let targets = RegressionTargetBlock {
1151 level: PredictionLevel::Sample,
1152 unit_ids: vec![sample_unit("sample:2"), sample_unit("sample:1")],
1153 values: vec![vec![5.0], vec![1.0]],
1154 target_names: vec!["y".to_string()],
1155 };
1156
1157 let report = score_regression_prediction_block(
1158 &predictions,
1159 &targets,
1160 &[
1161 RegressionMetricKind::Rmse,
1162 RegressionMetricKind::Mae,
1163 RegressionMetricKind::R2,
1164 ],
1165 )
1166 .unwrap();
1167
1168 assert_eq!(report.level, PredictionLevel::Sample);
1169 assert_close(report.metrics["rmse"], 1.0);
1170 assert_close(report.metrics["rmse:y"], 1.0);
1171 assert_close(report.metrics["mae"], 1.0);
1172 assert_close(report.metrics["r2"], 0.75);
1173 let candidate = regression_report_to_candidate_score("model:pls", report).unwrap();
1174 assert_eq!(candidate.metrics["rmse"], 1.0);
1175 assert_eq!(candidate.metadata["metric_level"], "sample");
1176 assert_eq!(candidate.metadata["producer_node"], "model:pls");
1177 assert_eq!(candidate.metadata["partition"], "validation");
1178 assert_eq!(candidate.metadata["prediction_id"], "pred:sample");
1179 assert_eq!(candidate.metadata["target_names"], serde_json::json!(["y"]));
1180 }
1181
1182 #[test]
1183 fn provider_adapter_preserves_display_target_names_with_spaces() {
1184 let predictions = PredictionBlock {
1185 prediction_id: None,
1186 producer_node: NodeId::new("model:pls").unwrap(),
1187 producer_port: Some("prediction".to_string()),
1188 partition: PredictionPartition::Validation,
1189 fold_id: None,
1190 sample_ids: vec![sid("sample:1"), sid("sample:2")],
1191 values: vec![vec![2.0], vec![4.0]],
1192 target_names: vec!["protein content".to_string()],
1193 };
1194 let targets = RegressionTargetBlock {
1195 level: PredictionLevel::Sample,
1196 unit_ids: vec![sample_unit("sample:1"), sample_unit("sample:2")],
1197 values: vec![vec![1.0], vec![5.0]],
1198 target_names: vec!["protein content".to_string()],
1199 };
1200
1201 let report = score_regression_prediction_block(
1202 &predictions,
1203 &targets,
1204 &[RegressionMetricKind::Rmse],
1205 )
1206 .unwrap();
1207 assert_close(report.metrics["rmse"], 1.0);
1208 assert_close(report.metrics["rmse:protein content"], 1.0);
1209 }
1210
1211 #[test]
1212 fn scores_target_and_group_prediction_blocks() {
1213 let predictions = AggregatedPredictionBlock {
1214 prediction_id: Some("pred:target".to_string()),
1215 producer_node: NodeId::new("model:pls").unwrap(),
1216 producer_port: None,
1217 partition: PredictionPartition::Validation,
1218 fold_id: None,
1219 level: PredictionLevel::Target,
1220 unit_ids: vec![target_unit("target:a"), target_unit("target:b")],
1221 values: vec![vec![1.0, 10.0], vec![3.0, 30.0]],
1222 target_names: vec!["y1".to_string(), "y2".to_string()],
1223 };
1224 let targets = RegressionTargetBlock {
1225 level: PredictionLevel::Target,
1226 unit_ids: vec![target_unit("target:b"), target_unit("target:a")],
1227 values: vec![vec![2.0, 28.0], vec![2.0, 12.0]],
1228 target_names: vec!["y1".to_string(), "y2".to_string()],
1229 };
1230 let report = score_regression_aggregated_block(
1231 &predictions,
1232 &targets,
1233 &[RegressionMetricKind::Mse, RegressionMetricKind::Rmse],
1234 )
1235 .unwrap();
1236
1237 assert_eq!(report.level, PredictionLevel::Target);
1238 assert_close(report.metrics["mse:y1"], 1.0);
1239 assert_close(report.metrics["mse:y2"], 4.0);
1240 assert_close(report.metrics["mse"], 2.5);
1241 assert_close(report.metrics["rmse:y1"], 1.0);
1242 assert_close(report.metrics["rmse:y2"], 2.0);
1243 assert_close(report.metrics["rmse"], 1.5);
1244
1245 let group_predictions = AggregatedPredictionBlock {
1246 prediction_id: Some("pred:group".to_string()),
1247 producer_node: NodeId::new("model:pls").unwrap(),
1248 producer_port: None,
1249 partition: PredictionPartition::Validation,
1250 fold_id: None,
1251 level: PredictionLevel::Group,
1252 unit_ids: vec![group_unit("group:a")],
1253 values: vec![vec![3.0]],
1254 target_names: vec!["y".to_string()],
1255 };
1256 let group_targets = RegressionTargetBlock {
1257 level: PredictionLevel::Group,
1258 unit_ids: vec![group_unit("group:a")],
1259 values: vec![vec![1.0]],
1260 target_names: vec!["y".to_string()],
1261 };
1262 let group_report = score_regression_aggregated_block(
1263 &group_predictions,
1264 &group_targets,
1265 &[RegressionMetricKind::Mae],
1266 )
1267 .unwrap();
1268 assert_eq!(group_report.level, PredictionLevel::Group);
1269 assert_close(group_report.metrics["mae"], 2.0);
1270 }
1271
1272 #[test]
1273 fn refuses_metric_alignment_and_contract_mismatches() {
1274 let predictions = AggregatedPredictionBlock {
1275 prediction_id: None,
1276 producer_node: NodeId::new("model:pls").unwrap(),
1277 producer_port: None,
1278 partition: PredictionPartition::Validation,
1279 fold_id: None,
1280 level: PredictionLevel::Target,
1281 unit_ids: vec![target_unit("target:a")],
1282 values: vec![vec![1.0]],
1283 target_names: vec!["y".to_string()],
1284 };
1285 let missing_target = RegressionTargetBlock {
1286 level: PredictionLevel::Target,
1287 unit_ids: vec![target_unit("target:b")],
1288 values: vec![vec![1.0]],
1289 target_names: vec!["y".to_string()],
1290 };
1291 assert!(score_regression_aggregated_block(
1292 &predictions,
1293 &missing_target,
1294 &[RegressionMetricKind::Rmse],
1295 )
1296 .is_err());
1297
1298 let wrong_level = RegressionTargetBlock {
1299 level: PredictionLevel::Group,
1300 unit_ids: vec![group_unit("group:a")],
1301 values: vec![vec![1.0]],
1302 target_names: vec!["y".to_string()],
1303 };
1304 assert!(score_regression_aggregated_block(
1305 &predictions,
1306 &wrong_level,
1307 &[RegressionMetricKind::Rmse],
1308 )
1309 .is_err());
1310
1311 assert!(score_regression_aggregated_block(&predictions, &missing_target, &[]).is_err());
1312 assert!(score_regression_aggregated_block(
1313 &predictions,
1314 &RegressionTargetBlock {
1315 level: PredictionLevel::Target,
1316 unit_ids: vec![target_unit("target:a")],
1317 values: vec![vec![1.0]],
1318 target_names: vec!["other".to_string()],
1319 },
1320 &[RegressionMetricKind::Rmse],
1321 )
1322 .is_err());
1323 assert!(score_regression_aggregated_block(
1324 &predictions,
1325 &RegressionTargetBlock {
1326 level: PredictionLevel::Target,
1327 unit_ids: vec![target_unit("target:a")],
1328 values: vec![vec![1.0]],
1329 target_names: vec!["y".to_string()],
1330 },
1331 &[RegressionMetricKind::Rmse, RegressionMetricKind::Rmse],
1332 )
1333 .is_err());
1334 }
1335
1336 #[test]
1337 fn refuses_duplicate_and_non_finite_sample_predictions() {
1338 let targets = RegressionTargetBlock {
1339 level: PredictionLevel::Sample,
1340 unit_ids: vec![sample_unit("sample:1")],
1341 values: vec![vec![1.0]],
1342 target_names: vec!["y".to_string()],
1343 };
1344 let mut predictions = PredictionBlock {
1345 prediction_id: None,
1346 producer_node: NodeId::new("model:pls").unwrap(),
1347 producer_port: None,
1348 partition: PredictionPartition::Validation,
1349 fold_id: None,
1350 sample_ids: vec![sid("sample:1")],
1351 values: vec![vec![f64::INFINITY]],
1352 target_names: vec!["y".to_string()],
1353 };
1354 assert!(score_regression_prediction_block(
1355 &predictions,
1356 &targets,
1357 &[RegressionMetricKind::Rmse],
1358 )
1359 .is_err());
1360
1361 predictions.values = vec![vec![1.0], vec![1.0]];
1362 predictions.sample_ids = vec![sid("sample:1"), sid("sample:1")];
1363 assert!(score_regression_prediction_block(
1364 &predictions,
1365 &targets,
1366 &[RegressionMetricKind::Rmse],
1367 )
1368 .is_err());
1369 }
1370
1371 #[test]
1372 fn constant_target_r2_is_finite_and_deterministic() {
1373 let targets = RegressionTargetBlock {
1374 level: PredictionLevel::Sample,
1375 unit_ids: vec![sample_unit("sample:1"), sample_unit("sample:2")],
1376 values: vec![vec![2.0], vec![2.0]],
1377 target_names: vec!["y".to_string()],
1378 };
1379 let exact_predictions = PredictionBlock {
1380 prediction_id: None,
1381 producer_node: NodeId::new("model:exact").unwrap(),
1382 producer_port: None,
1383 partition: PredictionPartition::Validation,
1384 fold_id: None,
1385 sample_ids: vec![sid("sample:1"), sid("sample:2")],
1386 values: vec![vec![2.0], vec![2.0]],
1387 target_names: vec!["y".to_string()],
1388 };
1389 let exact_report = score_regression_prediction_block(
1390 &exact_predictions,
1391 &targets,
1392 &[RegressionMetricKind::R2],
1393 )
1394 .unwrap();
1395 assert_close(exact_report.metrics["r2"], 1.0);
1396
1397 let off_predictions = PredictionBlock {
1398 values: vec![vec![2.0], vec![3.0]],
1399 ..exact_predictions
1400 };
1401 let off_report = score_regression_prediction_block(
1402 &off_predictions,
1403 &targets,
1404 &[RegressionMetricKind::R2],
1405 )
1406 .unwrap();
1407 assert_close(off_report.metrics["r2"], 0.0);
1408 }
1409
1410 fn score_report(
1411 partition: PredictionPartition,
1412 fold: Option<&str>,
1413 rmse: f64,
1414 ) -> RegressionMetricReport {
1415 RegressionMetricReport {
1416 prediction_id: None,
1417 producer_node: NodeId::new("model:compat.0").unwrap(),
1418 producer_port: None,
1419 variant_id: None,
1420 variant_label: None,
1421 partition,
1422 fold_id: fold.map(|value| FoldId::new(value).unwrap()),
1423 level: PredictionLevel::Sample,
1424 row_count: 10,
1425 target_width: 1,
1426 target_names: vec!["y".to_string()],
1427 metrics: BTreeMap::from([("rmse".to_string(), rmse), ("r2".to_string(), 0.5)]),
1428 }
1429 }
1430
1431 fn score_report_for_port(
1432 port: Option<&str>,
1433 partition: PredictionPartition,
1434 fold: Option<&str>,
1435 rmse: f64,
1436 ) -> RegressionMetricReport {
1437 RegressionMetricReport {
1438 producer_port: port.map(ToString::to_string),
1439 ..score_report(partition, fold, rmse)
1440 }
1441 }
1442
1443 #[test]
1444 fn score_set_round_trips_validates_and_rejects_duplicates() {
1445 let set = ScoreSet {
1446 schema_version: LEGACY_SCORE_SET_SCHEMA_VERSION,
1447 plan_id: "plan:demo".to_string(),
1448 selection_metric: Some("rmse".to_string()),
1449 reports: vec![
1450 score_report(PredictionPartition::Validation, Some("avg"), 18.75),
1451 score_report(PredictionPartition::Test, Some("final"), 13.28),
1452 ],
1453 };
1454 set.validate().unwrap();
1455
1456 let json = serde_json::to_string(&set).unwrap();
1458 let back: ScoreSet = serde_json::from_str(&json).unwrap();
1459 assert_eq!(back, set);
1460
1461 let parsed: ScoreSet =
1463 serde_json::from_value(serde_json::json!({"plan_id": "p", "reports": []})).unwrap();
1464 assert_eq!(parsed.schema_version, LEGACY_SCORE_SET_SCHEMA_VERSION);
1465
1466 let siblings = ScoreSet {
1468 schema_version: SCORE_SET_SCHEMA_VERSION,
1469 reports: vec![
1470 score_report_for_port(Some("pred"), PredictionPartition::Test, Some("final"), 1.0),
1471 score_report_for_port(Some("aux"), PredictionPartition::Test, Some("final"), 2.0),
1472 ],
1473 ..set.clone()
1474 };
1475 siblings.validate().unwrap();
1476
1477 let dup = ScoreSet {
1479 schema_version: SCORE_SET_SCHEMA_VERSION,
1480 reports: vec![
1481 score_report_for_port(Some("pred"), PredictionPartition::Test, Some("final"), 1.0),
1482 score_report_for_port(Some("pred"), PredictionPartition::Test, Some("final"), 2.0),
1483 ],
1484 ..set.clone()
1485 };
1486 assert!(dup.validate().is_err());
1487
1488 let legacy_with_port = ScoreSet {
1490 reports: vec![score_report_for_port(
1491 Some("pred"),
1492 PredictionPartition::Test,
1493 Some("final"),
1494 1.0,
1495 )],
1496 ..set.clone()
1497 };
1498 assert!(legacy_with_port.validate().is_err());
1499 let v2_without_port = ScoreSet {
1500 schema_version: SCORE_SET_SCHEMA_VERSION,
1501 reports: vec![score_report(PredictionPartition::Test, Some("final"), 1.0)],
1502 ..set.clone()
1503 };
1504 assert!(v2_without_port.validate().is_err());
1505
1506 let blank = ScoreSet {
1508 plan_id: " ".to_string(),
1509 reports: vec![score_report(PredictionPartition::Test, Some("final"), 1.0)],
1510 ..set
1511 };
1512 assert!(blank.validate().is_err());
1513 }
1514
1515 #[test]
1516 fn accuracy_and_balanced_accuracy_match_sklearn_on_imbalanced_classification() {
1517 let predictions = PredictionBlock {
1527 prediction_id: Some("pred:classif".to_string()),
1528 producer_node: NodeId::new("model:rf").unwrap(),
1529 producer_port: None,
1530 partition: PredictionPartition::Validation,
1531 fold_id: None,
1532 sample_ids: (0..10).map(|i| sid(&format!("s{i}"))).collect(),
1533 values: vec![
1534 vec![0.0],
1535 vec![0.0],
1536 vec![0.0],
1537 vec![0.0],
1538 vec![0.0],
1539 vec![0.0],
1540 vec![1.0],
1541 vec![0.0],
1542 vec![0.0],
1543 vec![0.0],
1544 ],
1545 target_names: vec!["y".to_string()],
1546 };
1547 let targets = RegressionTargetBlock {
1548 level: PredictionLevel::Sample,
1549 unit_ids: (0..10).map(|i| sample_unit(&format!("s{i}"))).collect(),
1550 values: vec![
1551 vec![0.0],
1552 vec![0.0],
1553 vec![0.0],
1554 vec![0.0],
1555 vec![0.0],
1556 vec![0.0],
1557 vec![1.0],
1558 vec![1.0],
1559 vec![2.0],
1560 vec![2.0],
1561 ],
1562 target_names: vec!["y".to_string()],
1563 };
1564
1565 let report = score_regression_prediction_block(
1566 &predictions,
1567 &targets,
1568 &[
1569 RegressionMetricKind::Accuracy,
1570 RegressionMetricKind::BalancedAccuracy,
1571 ],
1572 )
1573 .unwrap();
1574
1575 assert_close(report.metrics["accuracy"], 0.70);
1576 assert_close(report.metrics["balanced_accuracy"], 0.50);
1577 assert_eq!(
1579 RegressionMetricKind::BalancedAccuracy.objective(),
1580 MetricObjective::Maximize
1581 );
1582 }
1583
1584 #[test]
1585 fn cross_fold_balanced_accuracy_pools_oof_and_matches_sklearn() {
1586 let model = NodeId::new("model:rf").unwrap();
1594 let fold_block = |fold: &str, ids: &[usize], preds: &[f64]| PredictionBlock {
1595 prediction_id: Some(format!("pred:{fold}")),
1596 producer_node: model.clone(),
1597 producer_port: None,
1598 partition: PredictionPartition::Validation,
1599 fold_id: Some(FoldId::new(fold).unwrap()),
1600 sample_ids: ids.iter().map(|i| sid(&format!("s{i}"))).collect(),
1601 values: preds.iter().map(|p| vec![*p]).collect(),
1602 target_names: vec!["y".to_string()],
1603 };
1604 let target_record = |fold: &str, ids: &[usize], trues: &[f64]| RegressionTargetRecord {
1605 producer_node: model.clone(),
1606 producer_port: None,
1607 variant_id: None,
1608 partition: PredictionPartition::Validation,
1609 fold_id: Some(FoldId::new(fold).unwrap()),
1610 block: RegressionTargetBlock {
1611 level: PredictionLevel::Sample,
1612 unit_ids: ids.iter().map(|i| sample_unit(&format!("s{i}"))).collect(),
1613 values: trues.iter().map(|t| vec![*t]).collect(),
1614 target_names: vec!["y".to_string()],
1615 },
1616 };
1617
1618 let f0 = (0..5).collect::<Vec<_>>();
1624 let f1 = (5..10).collect::<Vec<_>>();
1625 let blocks = vec![
1626 fold_block("0", &f0, &[0.0, 0.0, 0.0, 1.0, 0.0]),
1627 fold_block("1", &f1, &[0.0, 0.0, 0.0, 0.0, 0.0]),
1628 ];
1629 let targets = vec![
1630 target_record("0", &f0, &[0.0, 0.0, 0.0, 1.0, 2.0]),
1631 target_record("1", &f1, &[0.0, 0.0, 0.0, 1.0, 2.0]),
1632 ];
1633
1634 let outcome = cross_fold_validation_reports(
1635 &blocks,
1636 &targets,
1637 &[
1638 RegressionMetricKind::Accuracy,
1639 RegressionMetricKind::BalancedAccuracy,
1640 ],
1641 FoldPartitionMode::Partition,
1642 )
1643 .unwrap();
1644
1645 assert_eq!(
1646 outcome.reports.len(),
1647 1,
1648 "one pooled `avg` report for the producer"
1649 );
1650 let avg = &outcome.reports[0];
1651 assert_eq!(avg.fold_id, Some(FoldId::new("avg").unwrap()));
1652 assert_eq!(avg.row_count, 10, "all OOF samples pooled exactly once");
1653 assert_close(avg.metrics["accuracy"], 0.70);
1654 assert_close(avg.metrics["balanced_accuracy"], 0.50);
1655
1656 assert_eq!(outcome.oof_averages.len(), 1, "one OOF average block");
1660 let oof = &outcome.oof_averages[0];
1661 assert_eq!(oof.predictions.partition, PredictionPartition::Validation);
1662 assert_eq!(oof.predictions.fold_id, Some(FoldId::new("avg").unwrap()));
1663 assert_eq!(oof.predictions.level, PredictionLevel::Sample);
1664 assert_eq!(oof.predictions.unit_ids.len(), 10);
1665 assert_eq!(oof.y_true.unit_ids, oof.predictions.unit_ids);
1666 assert_eq!(
1669 oof.predictions.values,
1670 vec![
1671 vec![0.0],
1672 vec![0.0],
1673 vec![0.0],
1674 vec![1.0],
1675 vec![0.0],
1676 vec![0.0],
1677 vec![0.0],
1678 vec![0.0],
1679 vec![0.0],
1680 vec![0.0],
1681 ]
1682 );
1683 assert_eq!(
1684 oof.y_true.values,
1685 vec![
1686 vec![0.0],
1687 vec![0.0],
1688 vec![0.0],
1689 vec![1.0],
1690 vec![2.0],
1691 vec![0.0],
1692 vec![0.0],
1693 vec![0.0],
1694 vec![1.0],
1695 vec![2.0],
1696 ]
1697 );
1698 }
1699}