"""Coverage (selective-prediction) curves. Question: if the most uncertain predictions are refused, how much does quality improve? A useful uncertainty measure buys accuracy when you discard. Method (see ai/ANALYSIS_PLAN.md 3.1) ----------------------------------- Rank samples by uncertainty ascending (most confident first, stable sort), retain the leading fraction c for each c in the coverage grid, and score accuracy and F1 on the retained subset. Repeated for every family x configuration x applicable uncertainty measure, at the clean baseline noise level. The ``single`` configuration evaluates each member on its own and averages the resulting curves across members (+/-1 std band), so it is a fair "one model" baseline rather than an arbitrary pick. A *flat* curve is a real result: it means the measure carries no information about correctness. A curve rising to the left is what a well-behaved uncertainty measure looks like. Outputs: ``coverage_accuracy.png``, ``coverage_f1.png`` (one panel per measure, one line per family x configuration) and ``coverage.json`` (every curve + AURC). """ import warnings from typing import Any, Dict, List import numpy as np import analysis.measures as ms import analysis.metrics as mt from analysis.context import CONFIG_SINGLE, AnalysisContext, Series from analysis.plotting import add_legend, panel_grid, save_figure, save_json, style_axes from analysis.registry import register def _series_curves( ctx: AnalysisContext, series: Series, measure: ms.Measure ) -> Dict[str, Any] | None: """Coverage curve(s) for one series and measure, aggregated over members.""" selected = ctx.clean_source(series.family) if selected is None: return None stem, ds = selected probs = ctx.probs(series.family, ds) # (M, S, C) if probs.shape[0] == 0: return None member_mi = ( ctx.member_stat("mutual_information", series.family, ds) if measure.needs_member_mi else None ) true_idx = ctx.true_index(ds) # ensemble -> one curve from all members; single -> one curve per member. if series.config == CONFIG_SINGLE: groups = [ (probs[m : m + 1], None if member_mi is None else member_mi[m : m + 1]) for m in range(probs.shape[0]) ] else: groups = [(probs, member_mi)] curves: List[mt.CoverageCurve] = [] for group_probs, group_mi in groups: uncertainty = ms.compute(measure, group_probs, member_mi=group_mi) pred = ms.predicted_class(group_probs) curves.append(mt.coverage_curve(uncertainty, pred, true_idx)) accuracy = np.array([c.accuracy for c in curves], dtype=float) f1 = np.array([c.f1 for c in curves], dtype=float) aurc_values = np.array([c.aurc for c in curves], dtype=float) # Low-coverage points can be NaN for every member (the min_samples guard); # averaging an all-NaN column legitimately yields NaN, so silence the notice. with warnings.catch_warnings(): warnings.simplefilter("ignore", RuntimeWarning) return { "source": stem, "coverage": curves[0].coverage, "n_retained": curves[0].n_retained, "n_curves": len(curves), "accuracy": np.nanmean(accuracy, axis=0).tolist(), "accuracy_std": np.nanstd(accuracy, axis=0).tolist(), "f1": np.nanmean(f1, axis=0).tolist(), "f1_std": np.nanstd(f1, axis=0).tolist(), "aurc": float(np.nanmean(aurc_values)), } def _plot_measure( ctx: AnalysisContext, measure_name: str, per_series: Dict[str, Dict[str, Any]], ) -> None: """One FIGURE per uncertainty measure: accuracy and F1 side by side. One file per measure keeps each panel to a handful of lines; a single grid of every measure was too dense to read. """ label = ms.MEASURES[measure_name].label fig, axes = panel_grid(2, n_cols=2, panel_size=(5.2, 3.8)) for ax, (metric, std_key, ylabel) in zip( axes, [ ("accuracy", "accuracy_std", "Accuracy"), ("f1", "f1_std", f"F1 ({mt.POSITIVE_CLASS})"), ], ): for payload in per_series.values(): series = payload["_series"] x = np.asarray(payload["coverage"], dtype=float) * 100.0 y = np.asarray(payload[metric], dtype=float) spread = np.asarray(payload[std_key], dtype=float) if payload["n_curves"] > 1: ax.fill_between( x, y - spread, y + spread, color=series.color, alpha=0.15, linewidth=0, ) ax.plot( x, y, color=series.color, linestyle=series.linestyle, linewidth=2, marker="o", markersize=3, label=series.label, ) ax.set_title(ylabel, fontsize=10) ax.set_xlabel("Coverage: most-confident samples retained (%)") ax.set_ylabel(ylabel) style_axes(ax) # Full dataset on the left, progressively more restricted to the right, # so the curve reads in the direction of "discard more". ax.invert_xaxis() fig.suptitle(f"Coverage curves — {label}") add_legend(fig, axes) png = save_figure(fig, ctx.out_dir, f"coverage_{measure_name}") ctx.log.info(f"coverage: wrote {png.name}.") @register( "coverage", title="Coverage curves: accuracy and F1 vs. retained fraction", ) def coverage(ctx: AnalysisContext) -> None: # results[measure][series_key] = curve payload results: Dict[str, Dict[str, Dict[str, Any]]] = { name: {} for name in ms.MEASURES } for series in ctx.series(): for measure in ms.measures_for(series.family, series.config): payload = _series_curves(ctx, series, measure) if payload is None: continue payload["_series"] = series results[measure.name][series.key] = payload if not any(results.values()): ctx.log.error("coverage: no usable evaluations.") return for measure_name, per_series in results.items(): if per_series: _plot_measure(ctx, measure_name, per_series) serializable = { measure: { key: {k: v for k, v in payload.items() if k != "_series"} for key, payload in per_series.items() } for measure, per_series in results.items() if per_series } save_json( { "positive_class": mt.POSITIVE_CLASS, "grid": list(mt.DEFAULT_COVERAGE_GRID), "min_samples": mt.MIN_COVERAGE_SAMPLES, "measures": serializable, }, ctx.out_dir, "coverage", ) for measure, per_series in serializable.items(): for key, payload in per_series.items(): ctx.log.info( f"coverage: {measure} / {key} — AURC={payload['aurc']:.4f} " f"(accuracy at full coverage {payload['accuracy'][-1]:.3f})." )