Skip to main content

dag_ml_core/
metrics.rs

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    /// Classification accuracy: fraction of predictions whose label matches the target (integer
28    /// label encoding, matched within 0.5). Meaningless on continuous regression targets (≈0) but
29    /// always emitted so the host can score classification natively without a separate code path.
30    Accuracy,
31    /// Balanced classification accuracy: the macro-average of per-class recall (mean over the
32    /// classes *present in `y_true`* of `correct_in_class / count_in_class`), matching scikit-learn's
33    /// `balanced_accuracy_score`. This is nirs4all's DEFAULT classification ranking metric (its
34    /// `_resolve_effective_metric` returns `balanced_accuracy` for a classification candidate), so it
35    /// must be emitted natively for the dag-ml engine to reproduce the legacy classification
36    /// `cv_best_score`. On a class-collapsed predictor it can be far below plain `accuracy`; on a
37    /// continuous regression target it is meaningless (≈ chance) but always emitted, like `accuracy`.
38    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    /// Resolve and validate the canonical metric/objective/output-kind matrix
73    /// shared by training request projection, native execution, and standalone
74    /// outcome verification.
75    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
164/// Mandatory, central *merge target-coverage* invariant — the single gate every merge reassembly
165/// handler (separation/concat, fusion, off-fold) must pass through before emitting (or declining to
166/// emit) a producer-level [`RegressionTargetBlock`] for a reassembled merge prediction. It closes
167/// audit R-P1-9: a merge that *should* be scored must never silently produce **no** score.
168///
169/// A merge reassembles its output's `y_true` by collecting the per-branch validation/off-fold target
170/// records into `by_sample_target`. The scoring path ([`super::runtime`]'s `apply_result_scoring`)
171/// pairs a prediction block 1:1 with a target block that covers *exactly* its samples — so a partial
172/// target block cannot be scored. There are exactly three legitimate outcomes:
173///
174/// 1. **No contributing branch emitted targets** (`by_sample_target` empty) — the merge is simply
175///    unscored. Returns `Ok(None)`; the caller emits no [`RegressionTargetBlock`]. This is the common
176///    "host never emitted `regression_targets`" case and stays a no-op (unchanged behavior).
177/// 2. **Every merged sample is covered** — emit the 1:1 target block. Returns `Ok(Some(block))` in the
178///    merge's declared `sample_id` order, ready to score.
179/// 3. **At least one branch emitted targets but coverage is INCOMPLETE** — previously the partial
180///    targets were silently dropped (`Vec::new()` → no score), so a merge that should have been scored
181///    silently vanished from selection. This is now a hard validation **ERROR**: once ANY branch
182///    contributes targets, the merge universe must be covered completely or the run fails loudly.
183///
184/// `merge_sample_ids` is the merge output's sample order (the universe to cover); `target_names` is the
185/// reassembled target name vector (already unified across branches by the caller). The map is consumed
186/// by `remove` on the success path so the caller need not clone it.
187pub 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    /// Variant this score belongs to — distinguishes per-variant candidates when a generated
237    /// campaign scores several variants, so native SELECT can pick the best. Skipped (None) for
238    /// single-variant runs, so existing fixtures/fingerprints are byte-identical.
239    #[serde(default, skip_serializing_if = "Option::is_none")]
240    pub variant_id: Option<VariantId>,
241    /// Cross-language CONTENT fingerprint (hex sha256) of the operator-variant this score belongs to:
242    /// the canonical form of the variant's LOWERED operator sub-sequence (Phase 5). The nirs4all host
243    /// recomputes the SAME bytes from its own operator-choice config, so it can map a per-variant
244    /// dag-ml report back to the config that produced it (replacing a brittle positional zip). Set
245    /// only for operator-SELECT reports (the choice fingerprint from `OperatorVariantModel`'s
246    /// `variant_labels`); skipped (None) for param-variant / single-variant runs, so existing
247    /// fixtures/fingerprints stay byte-identical.
248    #[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
416/// Current on-disk schema version for newly written [`ScoreSet`] documents.
417pub 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
425/// A persisted collection of per-block regression metric reports — the native, cross-language
426/// score record produced by a run (one report per `(producer_node, partition, fold_id, level)`).
427///
428/// Unlike the prediction-*value* cache (which is Validation-only and leakage-gated), scores are
429/// scalars derived from `y_true` and carry no feature data, so they are safe to persist for
430/// every partition (train / validation / test / final). This is the score half of "dag-ml owns
431/// prediction/score persistence natively" — the Python (or any host) `RunResult` reads these
432/// scalars by identity, with no recomputation.
433/// Identity of a score report within a [`ScoreSet`] — unique per producer port, variant,
434/// partition, fold and level.
435type 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    /// The metric SELECT optimized for this run (e.g. `"rmse"`), if a selection ran. Metadata only.
450    #[serde(default, skip_serializing_if = "Option::is_none")]
451    pub selection_metric: Option<String>,
452    pub reports: Vec<RegressionMetricReport>,
453}
454
455impl ScoreSet {
456    /// Validate the version, plan id, every report, and report-key uniqueness.
457    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
747/// Balanced classification accuracy for one target column: the macro-average of per-class recall over
748/// the integer labels present in `y_true`, matching scikit-learn's `balanced_accuracy_score`. Labels
749/// are matched the same way as [`RegressionMetricKind::Accuracy`] — a prediction counts for true class
750/// `c` when `|pred - c| < 0.5` — so the two metrics share one label-encoding convention. Returns the
751/// unweighted mean of `correct_in_class / count_in_class`; an empty target set yields `0.0` (the rows
752/// are non-empty here because the scoring path rejects zero-row blocks before this is reached).
753fn balanced_accuracy_for_target(
754    target_idx: usize,
755    predictions: &[&[f64]],
756    targets: &[&[f64]],
757) -> f64 {
758    // Group sample rows by their (rounded) true class label, preserving determinism via BTreeMap.
759    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/// A host-emitted `y_true` block tagged with the prediction it scores (producer/partition/fold), so
825/// the runtime can aggregate ground truth across folds to score cross-fold ensembles natively.
826#[derive(Clone, Debug, PartialEq)]
827pub struct RegressionTargetRecord {
828    pub producer_node: NodeId,
829    pub producer_port: Option<String>,
830    /// Variant that produced the scored block — lets the cross-fold OOF average be computed
831    /// per-variant (for native SELECT) without tagging every PredictionBlock with a variant.
832    pub variant_id: Option<VariantId>,
833    pub partition: PredictionPartition,
834    pub fold_id: Option<FoldId>,
835    pub block: RegressionTargetBlock,
836}
837
838/// Combine a producer's per-fold VALIDATION `y_true` into one block (dedup by unit id — a sample's
839/// ground truth is fold-independent), aligned to the producer's OOF samples.
840///
841/// Defense-in-depth (audit R-P0-1): records are grouped only by `producer_node`, but each carries a
842/// `variant_id`. A sample's ground truth is variant-independent, so the same unit seen again must
843/// carry the SAME `y_true`. If two records (e.g. from two variants sharing one context) disagree on a
844/// unit's target, the ground truth has been mixed and the combined block would silently score against
845/// a corrupted reference — that is refused rather than keeping whichever value happened to be first.
846fn 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/// The per-sample cross-fold OOF average of one producer, surfaced alongside its scalar report so the
890/// host can show each OOF sample's averaged prediction (nirs4all's `(validation, avg)` row y_pred),
891/// not only the pooled scalar. The block is keyed by `producer_node` / `partition = Validation` /
892/// `fold_id = "avg"` — identical to the scalar [`RegressionMetricReport`] this pairs with — and the
893/// `y_true` covers exactly the block's samples (same id set), so the host pairs them by id. This is
894/// REPORT-grade output: it carries no variant tag (the block has none; the variant is stamped on the
895/// report downstream) and never feeds a training/feature path, so OOF/leakage invariants are
896/// unaffected — it is purely the same averaged values the scalar was computed from, exposed per sample.
897#[derive(Clone, Debug, PartialEq)]
898pub struct OofAverageBlock {
899    pub predictions: AggregatedPredictionBlock,
900    pub y_true: RegressionTargetBlock,
901}
902
903/// The output of [`cross_fold_validation_reports`]: the scalar cross-fold OOF average reports (one per
904/// producer, `fold_id = "avg"`) plus — purely additively — the per-sample OOF average block + `y_true`
905/// each report was computed from. `reports` is byte-identical to the historical `Vec` return; callers
906/// that only need the scalars read `reports` and ignore `oof_averages`.
907#[derive(Clone, Debug, Default, PartialEq)]
908pub struct CrossFoldValidation {
909    pub reports: Vec<RegressionMetricReport>,
910    pub oof_averages: Vec<OofAverageBlock>,
911}
912
913/// Score the cross-fold OOF average per producer port: concatenate each `(producer_node,
914/// producer_port)` pair's per-fold VALIDATION predictions into one block and score it against the
915/// matching combined `y_true`. Yields one report per producer port with `fold_id = "avg"` —
916/// nirs4all's `cv_best_score` row — plus, additively, the per-sample OOF average block + `y_true`
917/// each report was computed from (so the host can fill the `(validation, avg)` row's per-sample
918/// y_pred, not only the scalar). The per-fold join is identity-keyed; producer ports with a single
919/// fold are skipped (nothing to ensemble).
920///
921/// `partition_mode` mirrors the campaign [`FoldPartitionMode`]:
922/// under `Partition` (KFold) the per-producer OOF must be unique (each sample scored exactly once);
923/// under `Resampled` (ShuffleSplit / repeated CV) a sample may appear in several folds — those
924/// predictions are averaged by [`reduce_predictions_across_folds`] — so the across-fold uniqueness
925/// gate is relaxed accordingly.
926pub 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        // Mandatory OOF coverage gate (spec rule 3), mode-aware. Under `Partition` the producer's
952        // per-fold validation blocks must be UNIQUE — exactly one validation prediction per sample; a
953        // sample appearing in two blocks would be a duplicated fold or — since `PredictionBlock` carries
954        // no variant tag — two variants' OOF in a shared context (audit R-P0-1), and is refused (the
955        // scoring-path analogue of the runtime merge handler's "mixes several variants" guard, so
956        // cross-variant CV scores can NEVER mix here). Under `Resampled` (ShuffleSplit / repeated CV) a
957        // sample is legitimately validated in several folds and its predictions are averaged by
958        // `reduce_predictions_across_folds` below, so across-fold multiplicity is allowed; the per-block
959        // within-fold uniqueness still holds via `validate_content`.
960        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            // No y_true was emitted for this producer (e.g. mock controllers) — nothing to score.
965            continue;
966        }
967        let average = reduce_predictions_across_folds(blocks, None, "avg")?;
968        // The scalar report is computed from `average` EXACTLY as before — byte-identical. The
969        // additive per-sample surface below reuses the SAME `average` values and the SAME `targets`,
970        // so it cannot perturb any score or `row_count`.
971        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
982/// Build the per-sample OOF average surface (block + `y_true`) from the SAME `average`
983/// [`PredictionBlock`] and combined `targets` the scalar report was computed from. The block is the
984/// sample-level lift of `average` (its sample ids become `Sample` unit ids, values unchanged), keyed
985/// identically (producer / `Validation` / `avg`). The `y_true` is `targets` realigned to the block's
986/// sample order so a host pairs y_pred ↔ y_true per sample without re-sorting; every `average` sample
987/// has a y_true row because [`combine_validation_targets`] pools every per-fold validation record and
988/// the OOF coverage gate guarantees each averaged sample was validated.
989fn 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        // No contributing branch emitted targets: the merge is legitimately unscored.
1076        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        // Emitted in the merge's declared sample order, not map order.
1106        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        // R-P1-9: one branch contributed a target (s1) but the merge universe also
1119        // covers s2, which has no y_true row. Previously this silently dropped the
1120        // targets (no score); it must now be a hard validation error.
1121        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        // JSON round-trip is lossless.
1457        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        // schema_version defaults when omitted (forward-compatible read).
1462        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        // Sibling ports of the same node are distinct score identities.
1467        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        // Duplicate (producer_node, producer_port, partition, fold_id, level) is rejected.
1478        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        // Families are all-or-nothing: V1 forbids ports; V2 requires them.
1489        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        // Empty plan_id is rejected.
1507        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        // #60 root-cause lock: dag-ml emits BOTH plain `accuracy` and `balanced_accuracy`. nirs4all's
1518        // DEFAULT classification ranking metric is balanced_accuracy (its `_resolve_effective_metric`),
1519        // so the legacy `cv_best_score` for a classification sweep is balanced_accuracy — NOT plain
1520        // accuracy. A class-collapsed predictor on imbalanced data makes the two diverge sharply (the
1521        // 0.32-vs-0.16 STOP report). Ground truth here is scikit-learn on the same labels:
1522        //   y    = [0,0,0,0,0,0, 1,1, 2,2]  (majority class 0)
1523        //   pred = [0,0,0,0,0,0, 1,0, 0,0]  (collapses wrong rows to class 0)
1524        //   accuracy_score          = 7/10 = 0.70
1525        //   balanced_accuracy_score = mean(recall(c0)=6/6, recall(c1)=1/2, recall(c2)=0/2) = 0.50
1526        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        // Both maximize — the host SELECT ranks them in the same direction.
1578        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        // The real #60 path: per-fold VALIDATION OOF blocks pooled into the `avg` report (nirs4all's
1587        // `cv_best_score` row), scored against the combined y_true. Two disjoint KFold folds carry the
1588        // same imbalanced class-collapse as the per-block lock above, so the POOLED OOF accuracy is
1589        // 0.70 and pooled balanced_accuracy is 0.50 — proving the cross-fold reduction + metric agree
1590        // with scikit-learn's `accuracy_score` / `balanced_accuracy_score` on the same labels, and that
1591        // dag-ml's `accuracy` (0.70) was never wrong: it is simply a DIFFERENT metric from the legacy's
1592        // default `balanced_accuracy` (0.50).
1593        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        // Fold 0 (samples 0..5): preds [0,0,0,1,0] vs true [0,0,0,1,2]
1619        // Fold 1 (samples 5..10): preds [0,0,0,0,0] vs true [0,0,0,1,2]
1620        // Pooled: true=[0,0,0,1,2,0,0,0,1,2], pred=[0,0,0,1,0,0,0,0,0,0]
1621        //   accuracy          = 7/10 = 0.70
1622        //   balanced_accuracy = mean(recall0=6/6, recall1=1/2, recall2=0/2) = 0.50
1623        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        // Additive per-sample OOF average surface: one block per scored producer, keyed identically to
1657        // the scalar report (producer / Validation / avg), with the SAME pooled values and one y_true
1658        // row per averaged sample (same id set), realigned to the block's sample order.
1659        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        // The pooled per-sample preds are the across-fold mean (each KFold sample validated once):
1667        // [0,0,0,1,0] ++ [0,0,0,0,0] against y_true [0,0,0,1,2] ++ [0,0,0,1,2].
1668        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}