"""Performance and uncertainty vs. input noise. Question: how do accuracy and each uncertainty measure respond as Gaussian input noise increases? This is the core robustness result. Method (see ai/ANALYSIS_PLAN.md 3.2) ----------------------------------- Uses a noise-sweep evaluation (``noisy_*`` / ``combined_noisy_*``). For every noise level sigma, computes accuracy, F1 and the mean (+/-1 std across samples) of each applicable uncertainty measure, for every family x configuration. The ``single`` configuration averages over members. The scientifically interesting comparison is the *rate* of change: the epistemic terms (``ensemble_mutual_information``, ``mean_mutual_information``) should rise faster than the aleatoric term (``mean_predictive_entropy``) as inputs move out of distribution. If they do not, the uncertainty is not tracking what it should. Outputs: ``noise_performance.png`` (accuracy and F1 panels, then one panel per uncertainty measure, all against sigma) and ``noise_performance.json``. """ from typing import Any, Dict, List, Tuple import numpy as np import analysis.measures as ms import analysis.metrics as mt from analysis.context import CONFIG_SINGLE, AnalysisContext, Series from analysis.sources import noise_axis_label from analysis.plotting import add_legend, panel_grid, save_figure, save_json, style_axes from analysis.registry import register def _member_groups( probs: np.ndarray, member_mi: np.ndarray | None, config: str ) -> List[Tuple[np.ndarray, np.ndarray | None]]: """Split into the units scored independently: 1 ensemble, or M singles.""" if config == CONFIG_SINGLE: return [ (probs[m : m + 1], None if member_mi is None else member_mi[m : m + 1]) for m in range(probs.shape[0]) ] return [(probs, member_mi)] def _series_response( ctx: AnalysisContext, series: Series ) -> Dict[str, Any] | None: """Accuracy, F1 and every applicable measure across the noise sweep.""" selected = ctx.noise_source(series.family) if selected is None: return None stem, ds = selected levels = ctx.noise_levels(ds) true_idx = ctx.true_index(ds) applicable = ms.measures_for(series.family, series.config) accuracy: List[float] = [] accuracy_std: List[float] = [] f1: List[float] = [] f1_std: List[float] = [] measure_mean: Dict[str, List[float]] = {m.name: [] for m in applicable} measure_std: Dict[str, List[float]] = {m.name: [] for m in applicable} needs_mi = any(m.needs_member_mi for m in applicable) for level_index in range(len(levels)): probs = ctx.probs(series.family, ds, level_index) member_mi = ( ctx.member_stat("mutual_information", series.family, ds, level_index) if needs_mi else None ) groups = _member_groups(probs, member_mi, series.config) # Performance: score each unit, then average over units. unit_acc = [ mt.accuracy(ms.predicted_class(g), true_idx) for g, _ in groups ] unit_f1 = [mt.f1_score(ms.predicted_class(g), true_idx) for g, _ in groups] accuracy.append(float(np.nanmean(unit_acc))) accuracy_std.append(float(np.nanstd(unit_acc))) f1.append(float(np.nanmean(unit_f1))) f1_std.append(float(np.nanstd(unit_f1))) # Uncertainty: mean over samples for each unit, then averaged over units; # the band is the spread across samples (pooled over units). for measure in applicable: per_sample = np.concatenate( [ms.compute(measure, g, member_mi=gm) for g, gm in groups] ) measure_mean[measure.name].append(float(np.mean(per_sample))) measure_std[measure.name].append(float(np.std(per_sample))) return { "source": stem, "noise_axis": noise_axis_label(ds), "noise_levels": levels.tolist(), "accuracy": accuracy, "accuracy_std": accuracy_std, "f1": f1, "f1_std": f1_std, "measures": { name: {"mean": measure_mean[name], "std": measure_std[name]} for name in measure_mean }, } @register( "noise_performance", title="Accuracy and uncertainty vs. input noise", requires_noise=True, ) def noise_performance(ctx: AnalysisContext) -> None: responses: Dict[str, Dict[str, Any]] = {} for series in ctx.series(ctx.noise_families): payload = _series_response(ctx, series) if payload is not None: payload["_series"] = series responses[series.key] = payload if not responses: ctx.log.error("noise_performance: no noise-sweep evaluations available.") return # Panels: accuracy, F1, then every measure that any series produced. measure_panels = [ name for name in ms.MEASURES if any(name in p["measures"] for p in responses.values()) ] panel_keys = ["accuracy", "f1"] + measure_panels fig, axes = panel_grid(len(panel_keys), n_cols=2) for ax, key in zip(axes, panel_keys): is_performance = key in ("accuracy", "f1") for payload in responses.values(): series: Series = payload["_series"] x = np.asarray(payload["noise_levels"], dtype=float) if is_performance: y = np.asarray(payload[key], dtype=float) spread = np.asarray(payload[f"{key}_std"], dtype=float) show_band = series.config == CONFIG_SINGLE else: if key not in payload["measures"]: continue y = np.asarray(payload["measures"][key]["mean"], dtype=float) spread = np.asarray(payload["measures"][key]["std"], dtype=float) show_band = True if show_band: ax.fill_between( x, y - spread, y + spread, color=series.color, alpha=0.13, linewidth=0, ) ax.plot( x, y, color=series.color, linestyle=series.linestyle, linewidth=2, marker="o", markersize=4, label=series.label, ) if key == "accuracy": ax.set_title("Accuracy", fontsize=10) ax.set_ylabel("Accuracy") elif key == "f1": ax.set_title(f"F1 ({mt.POSITIVE_CLASS})", fontsize=10) ax.set_ylabel("F1") else: ax.set_title(ms.MEASURES[key].label, fontsize=10) ax.set_ylabel("Uncertainty (nats)") ax.set_xlabel(next(iter(responses.values()))["noise_axis"]) style_axes(ax) fig.suptitle("Performance and uncertainty vs. input noise") add_legend(fig, axes) png = save_figure(fig, ctx.out_dir, "noise_performance") save_json( { "positive_class": mt.POSITIVE_CLASS, "series": { key: {k: v for k, v in payload.items() if k != "_series"} for key, payload in responses.items() }, }, ctx.out_dir, "noise_performance", ) for key, payload in responses.items(): levels = payload["noise_levels"] acc = payload["accuracy"] ctx.log.info( f"noise_performance: {key} — accuracy {acc[0]:.3f} at σ={levels[0]:g} " f"→ {acc[-1]:.3f} at σ={levels[-1]:g}." ) ctx.log.info(f"noise_performance: wrote {png.name} ({len(panel_keys)} panels).")