| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182 |
- """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(
- ctx: AnalysisContext,
- results: Dict[str, Dict[str, Dict[str, Any]]],
- metric: str,
- std_key: str,
- ylabel: str,
- stem: str,
- ) -> None:
- """One panel per measure; one line per series."""
- panels = [name for name, per_series in results.items() if per_series]
- if not panels:
- return
- fig, axes = panel_grid(len(panels), n_cols=2)
- for ax, measure_name in zip(axes, panels):
- for series_key, payload in results[measure_name].items():
- 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(ms.MEASURES[measure_name].label, fontsize=10)
- ax.set_xlabel("Coverage: most-confident samples retained (%)")
- ax.set_ylabel(ylabel)
- style_axes(ax)
- fig.suptitle(f"Coverage curves — {ylabel} vs. retained fraction")
- add_legend(fig, axes)
- png = save_figure(fig, ctx.out_dir, stem)
- ctx.log.info(f"coverage: wrote {png.name} ({len(panels)} measure panels).")
- @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
- _plot(ctx, results, "accuracy", "accuracy_std", "Accuracy", "coverage_accuracy")
- _plot(ctx, results, "f1", "f1_std", "F1 (AD)", "coverage_f1")
- 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})."
- )
|