noise_correlation.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232
  1. """Uncertainty vs. accuracy across noise conditions, with a fitted curve.
  2. Question: is uncertainty a *transferable* predictor of error -- does a given
  3. uncertainty value imply the same accuracy regardless of the noise condition?
  4. Method (see ai/ANALYSIS_PLAN.md 3.3)
  5. -----------------------------------
  6. For each noise level sigma, split its samples into ``N_BINS`` equal-count
  7. (quantile) uncertainty bins and plot one point per bin: mean uncertainty in the
  8. bin (x) against accuracy in the bin (y), coloured by sigma. This gives
  9. sigma x bins points -- enough to fit -- rather than the single point per sigma a
  10. naive reading would produce.
  11. Fits ``acc = a + b * exp(-k * u)`` (bounded, monotone) with a linear fallback,
  12. and reports R^2 plus Spearman rho. Rho is the robust headline because it makes no
  13. assumption about the functional form.
  14. Interpretation
  15. --------------
  16. If points from every sigma collapse onto one curve, the measure is calibrated
  17. *across conditions*: uncertainty alone predicts accuracy. If they separate by
  18. sigma, the measure is condition-dependent and cannot be used as a standalone
  19. confidence signal -- an important negative result if it appears.
  20. Outputs: ``noise_correlation_<measure>.png`` (one panel per family x
  21. configuration) and ``noise_correlation.json`` (fit parameters, R^2, rho).
  22. """
  23. from typing import Any, Dict, List
  24. import numpy as np
  25. from matplotlib import cm
  26. from matplotlib.colors import Normalize
  27. import analysis.measures as ms
  28. import analysis.metrics as mt
  29. from analysis.context import CONFIG_SINGLE, AnalysisContext, Series
  30. from analysis.sources import noise_axis_label
  31. from analysis.plotting import panel_grid, save_figure, save_json, style_axes
  32. from analysis.registry import register
  33. #: Equal-count uncertainty bins per noise level.
  34. N_BINS = 10
  35. #: Minimum samples in a bin for its accuracy to be trustworthy.
  36. MIN_BIN_SAMPLES = 5
  37. def _series_points(
  38. ctx: AnalysisContext, series: Series, measure: ms.Measure
  39. ) -> Dict[str, Any] | None:
  40. """Binned (uncertainty, accuracy, sigma) points plus the fit, for one series."""
  41. selected = ctx.noise_source(series.family)
  42. if selected is None:
  43. return None
  44. stem, ds = selected
  45. levels = ctx.noise_levels(ds)
  46. true_idx = ctx.true_index(ds)
  47. x_values: List[float] = []
  48. y_values: List[float] = []
  49. sigma_values: List[float] = []
  50. for level_index, sigma in enumerate(levels):
  51. probs = ctx.probs(series.family, ds, level_index)
  52. member_mi = (
  53. ctx.member_stat("mutual_information", series.family, ds, level_index)
  54. if measure.needs_member_mi
  55. else None
  56. )
  57. # `single` pools every (member, sample) observation at this noise level,
  58. # so it yields the same number of points as `ensemble` while using all
  59. # members' predictions.
  60. if series.config == CONFIG_SINGLE:
  61. uncertainty = np.concatenate(
  62. [
  63. ms.compute(
  64. measure,
  65. probs[m : m + 1],
  66. member_mi=None if member_mi is None else member_mi[m : m + 1],
  67. )
  68. for m in range(probs.shape[0])
  69. ]
  70. )
  71. correct = np.concatenate(
  72. [
  73. ms.predicted_class(probs[m : m + 1]) == true_idx
  74. for m in range(probs.shape[0])
  75. ]
  76. )
  77. else:
  78. uncertainty = ms.compute(measure, probs, member_mi=member_mi)
  79. correct = ms.predicted_class(probs) == true_idx
  80. for indices in mt.quantile_bins(uncertainty, N_BINS):
  81. if indices.size < MIN_BIN_SAMPLES:
  82. continue
  83. x_values.append(float(np.mean(uncertainty[indices])))
  84. y_values.append(float(np.mean(correct[indices])))
  85. sigma_values.append(float(sigma))
  86. if len(x_values) < 3:
  87. return None
  88. fit = mt.fit_accuracy_vs_uncertainty(
  89. np.asarray(x_values), np.asarray(y_values)
  90. )
  91. return {
  92. "source": stem,
  93. "noise_axis": noise_axis_label(ds),
  94. "uncertainty": x_values,
  95. "accuracy": y_values,
  96. "noise_level": sigma_values,
  97. "n_bins": N_BINS,
  98. "fit": {
  99. "model": fit.model,
  100. "params": fit.params,
  101. "r_squared": fit.r_squared,
  102. "spearman": fit.spearman,
  103. },
  104. "_fit_line": (fit.x_fit, fit.y_fit),
  105. }
  106. def _plot_measure(
  107. ctx: AnalysisContext,
  108. measure: ms.Measure,
  109. per_series: Dict[str, Dict[str, Any]],
  110. ) -> None:
  111. """One panel per series; points coloured by noise level, with the fit drawn."""
  112. keys = list(per_series)
  113. fig, axes = panel_grid(len(keys), n_cols=2, panel_size=(5.0, 3.8))
  114. all_sigmas = sorted(
  115. {s for payload in per_series.values() for s in payload["noise_level"]}
  116. )
  117. norm = Normalize(vmin=min(all_sigmas), vmax=max(all_sigmas))
  118. colormap = cm.viridis
  119. scatter = None
  120. for ax, key in zip(axes, keys):
  121. payload = per_series[key]
  122. series: Series = payload["_series"]
  123. scatter = ax.scatter(
  124. payload["uncertainty"],
  125. payload["accuracy"],
  126. c=payload["noise_level"],
  127. cmap=colormap,
  128. norm=norm,
  129. s=26,
  130. edgecolor="white",
  131. linewidth=0.4,
  132. zorder=3,
  133. )
  134. x_fit, y_fit = payload["_fit_line"]
  135. if x_fit:
  136. ax.plot(x_fit, y_fit, color="#444444", linewidth=1.6, zorder=2)
  137. fit = payload["fit"]
  138. ax.annotate(
  139. f"{fit['model']} R²={fit['r_squared']:.3f}\nρ={fit['spearman']:.3f}",
  140. xy=(0.97, 0.95),
  141. xycoords="axes fraction",
  142. ha="right",
  143. va="top",
  144. fontsize=8,
  145. color="#333333",
  146. )
  147. ax.set_title(series.label, fontsize=10)
  148. ax.set_xlabel(f"{measure.label}")
  149. ax.set_ylabel("Accuracy (per bin)")
  150. style_axes(ax)
  151. if scatter is not None:
  152. label = next(iter(per_series.values()))["noise_axis"]
  153. fig.colorbar(scatter, ax=axes, label=label, shrink=0.85)
  154. fig.suptitle(f"Accuracy vs. {measure.label.lower()} across noise levels")
  155. stem = f"noise_correlation_{measure.name}"
  156. png = save_figure(fig, ctx.out_dir, stem)
  157. ctx.log.info(f"noise_correlation: wrote {png.name}.")
  158. @register(
  159. "noise_correlation",
  160. title="Uncertainty vs. accuracy across noise levels (with fit)",
  161. requires_noise=True,
  162. )
  163. def noise_correlation(ctx: AnalysisContext) -> None:
  164. results: Dict[str, Dict[str, Dict[str, Any]]] = {}
  165. for series in ctx.series(ctx.noise_families):
  166. for measure in ms.measures_for(series.family, series.config):
  167. payload = _series_points(ctx, series, measure)
  168. if payload is None:
  169. continue
  170. payload["_series"] = series
  171. results.setdefault(measure.name, {})[series.key] = payload
  172. if not results:
  173. ctx.log.error("noise_correlation: no noise-sweep evaluations available.")
  174. return
  175. for measure_name, per_series in results.items():
  176. _plot_measure(ctx, ms.MEASURES[measure_name], per_series)
  177. for key, payload in per_series.items():
  178. fit = payload["fit"]
  179. ctx.log.info(
  180. f"noise_correlation: {measure_name} / {key} — {fit['model']} fit "
  181. f"R²={fit['r_squared']:.3f}, Spearman ρ={fit['spearman']:.3f}."
  182. )
  183. save_json(
  184. {
  185. "n_bins": N_BINS,
  186. "min_bin_samples": MIN_BIN_SAMPLES,
  187. "measures": {
  188. measure: {
  189. key: {
  190. k: v
  191. for k, v in payload.items()
  192. if k not in ("_series", "_fit_line")
  193. }
  194. for key, payload in per_series.items()
  195. }
  196. for measure, per_series in results.items()
  197. },
  198. },
  199. ctx.out_dir,
  200. "noise_correlation",
  201. )