NeerajCodz's picture
Complete reviewer 2026-09 revision
8b37c3f
Raw History Blame Contribute Delete
4.6 kB
"""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