File size: 4,603 Bytes
8b37c3f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
"""Battery-grouped validation, clustered uncertainty, and paired testing."""

from __future__ import annotations

from collections.abc import Iterator

import numpy as np
import pandas as pd
from scipy.stats import wilcoxon
from sklearn.model_selection import GroupKFold, GroupShuffleSplit

from src.evaluation.metrics import regression_metrics


def grouped_train_val_test_folds(
    groups: np.ndarray | pd.Series,
    *,
    n_splits: int = 5,
    validation_fraction: float = 0.2,
    random_state: int = 42,
) -> Iterator[tuple[np.ndarray, np.ndarray, np.ndarray]]:
    """Yield disjoint battery-grouped train, validation, and test indices."""
    groups = np.asarray(groups)
    indices = np.arange(groups.size)
    outer = GroupKFold(n_splits=n_splits, shuffle=True, random_state=random_state)
    for train_val_idx, test_idx in outer.split(indices, groups=groups):
        splitter = GroupShuffleSplit(
            n_splits=1,
            test_size=validation_fraction,
            random_state=random_state,
        )
        rel_train, rel_val = next(
            splitter.split(train_val_idx, groups=groups[train_val_idx])
        )
        train_idx = train_val_idx[rel_train]
        val_idx = train_val_idx[rel_val]
        assert not (
            set(groups[train_idx]) & set(groups[val_idx])
            or set(groups[train_idx]) & set(groups[test_idx])
            or set(groups[val_idx]) & set(groups[test_idx])
        )
        yield train_idx, val_idx, test_idx


def battery_cluster_bootstrap(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    battery_ids: np.ndarray | pd.Series,
    *,
    n_predictors: int | None = None,
    n_bootstrap: int = 10_000,
    confidence: float = 0.95,
    random_state: int = 42,
) -> pd.DataFrame:
    """Return battery-cluster bootstrap confidence intervals for all metrics."""
    y_true = np.asarray(y_true)
    y_pred = np.asarray(y_pred)
    battery_ids = np.asarray(battery_ids)
    unique = np.unique(battery_ids)
    if unique.size < 2:
        raise ValueError("At least two batteries are required for clustered bootstrap")
    rng = np.random.default_rng(random_state)
    draws: list[dict[str, float]] = []
    group_indices = {group: np.flatnonzero(battery_ids == group) for group in unique}
    for _ in range(n_bootstrap):
        sampled = rng.choice(unique, size=unique.size, replace=True)
        idx = np.concatenate([group_indices[group] for group in sampled])
        draws.append(regression_metrics(y_true[idx], y_pred[idx], n_predictors=n_predictors))
    frame = pd.DataFrame(draws)
    alpha = (1.0 - confidence) / 2.0
    point = regression_metrics(y_true, y_pred, n_predictors=n_predictors)
    return pd.DataFrame({
        "metric": list(point),
        "estimate": [point[key] for key in point],
        "ci_low": [frame[key].quantile(alpha) for key in point],
        "ci_high": [frame[key].quantile(1.0 - alpha) for key in point],
    })


def holm_adjust(p_values: np.ndarray | list[float]) -> np.ndarray:
    """Holm family-wise error correction in original hypothesis order."""
    p = np.asarray(p_values, dtype=float)
    order = np.argsort(p)
    adjusted_sorted = np.maximum.accumulate((p.size - np.arange(p.size)) * p[order])
    adjusted = np.empty_like(p)
    adjusted[order] = np.clip(adjusted_sorted, 0.0, 1.0)
    return adjusted


def paired_battery_wilcoxon(
    scores: pd.DataFrame,
    *,
    battery_col: str = "battery_id",
    model_col: str = "model",
    score_col: str = "mae",
    reference_model: str,
) -> pd.DataFrame:
    """Compare each model with a reference using paired per-battery scores."""
    pivot = scores.pivot(index=battery_col, columns=model_col, values=score_col)
    if reference_model not in pivot:
        raise KeyError(f"Reference model not found: {reference_model}")
    rows = []
    for model in pivot.columns:
        if model == reference_model:
            continue
        pair = pivot[[reference_model, model]].dropna()
        if len(pair) < 2:
            continue
        statistic, p_value = wilcoxon(
            pair[model], pair[reference_model], zero_method="zsplit"
        )
        rows.append({
            "reference_model": reference_model,
            "model": model,
            "n_batteries": len(pair),
            "median_mae_difference": float(np.median(pair[model] - pair[reference_model])),
            "statistic": float(statistic),
            "p_value": float(p_value),
        })
    result = pd.DataFrame(rows)
    if not result.empty:
        result["p_value_holm"] = holm_adjust(result["p_value"].to_numpy())
    return result