model_report.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109
  1. """Per-model summary report (PDF, via Typst).
  2. A compact table of every ensemble member's accuracy and mean uncertainty, one
  3. section per model family, plus the ensemble headline. Data is injected into the
  4. Typst template as a JSON string through ``sys.inputs``.
  5. """
  6. import datetime
  7. import json
  8. import pathlib as pl
  9. from typing import Any, Dict, List
  10. import numpy as np
  11. import typst
  12. from analysis.sources import FAMILY_LABELS, baseline_noise_index
  13. from analysis.context import AnalysisContext
  14. from analysis.plotting import data_dir
  15. from analysis.registry import register
  16. _TEMPLATE = pl.Path(__file__).resolve().parent.parent / "templates" / "model_report.typ"
  17. def _family_block(ctx: AnalysisContext, family: str) -> Dict[str, Any] | None:
  18. """Per-member and ensemble metrics for one family at the clean baseline."""
  19. selected = ctx.clean_source(family)
  20. if selected is None:
  21. return None
  22. stem, ds = selected
  23. baseline = baseline_noise_index(ds)
  24. positions = ctx.member_positions(family, ds)
  25. correct = np.asarray(ds["correct"].isel(noise_level=baseline).values, dtype=float)
  26. entropy = ctx.member_stat("predictive_entropy", family, ds, baseline)
  27. mut_info = ctx.member_stat("mutual_information", family, ds, baseline)
  28. labels = np.atleast_1d(ds["model"].values)
  29. models: List[Dict[str, Any]] = []
  30. accuracies: List[float] = []
  31. for row, position in enumerate(positions):
  32. acc = float(np.mean(correct[position]))
  33. accuracies.append(acc)
  34. models.append(
  35. {
  36. "index": int(labels[position]) + 1, # 1-based for display
  37. "accuracy": acc,
  38. "mean_entropy": float(np.mean(entropy[row])),
  39. "mean_mi": float(np.mean(mut_info[row])),
  40. }
  41. )
  42. acc_arr = np.asarray(accuracies, dtype=float)
  43. return {
  44. "name": FAMILY_LABELS.get(family, family),
  45. "kind": family,
  46. "source": stem,
  47. "split": str(ds.attrs.get("split", "unknown")),
  48. "noise_sigma": float(ctx.noise_levels(ds)[baseline]),
  49. "n_models": len(models),
  50. "n_samples": int(ds.sizes["sample"]),
  51. "n_mc": int(ds.attrs.get("n_mc", 1)),
  52. "accuracy_mean": float(acc_arr.mean()) if acc_arr.size else 0.0,
  53. "accuracy_std": float(acc_arr.std()) if acc_arr.size else 0.0,
  54. "accuracy_best": float(acc_arr.max()) if acc_arr.size else 0.0,
  55. "models": models,
  56. }
  57. @register(
  58. "model_report",
  59. title="Per-model accuracy and uncertainty summary (PDF)",
  60. )
  61. def model_report(ctx: AnalysisContext) -> None:
  62. families = [
  63. block
  64. for block in (_family_block(ctx, family) for family in ctx.families)
  65. if block is not None
  66. ]
  67. if not families:
  68. ctx.log.error("model_report: no model families found in evaluations.")
  69. return
  70. meta = next(iter(ctx.datasets.values())).attrs
  71. seed = meta.get("seed")
  72. payload = {
  73. "title": "Model Evaluation Report",
  74. "generated": datetime.datetime.now(datetime.timezone.utc).strftime(
  75. "%Y-%m-%d %H:%M UTC"
  76. ),
  77. "work_dir": str(ctx.out_dir.parent),
  78. "schema_version": str(meta.get("schema_version", "?")),
  79. "seed": str(seed) if seed is not None else "n/a",
  80. "git_commit": str(meta.get("git_commit", "unknown"))[:10],
  81. "families": families,
  82. }
  83. payload_json = json.dumps(payload)
  84. # The template parses this with `json(bytes(sys.inputs.at("data")))`.
  85. pdf_bytes = typst.compile(
  86. str(_TEMPLATE), sys_inputs={"data": payload_json}, format="pdf"
  87. )
  88. (ctx.out_dir / "model_report.pdf").write_bytes(pdf_bytes)
  89. (data_dir(ctx.out_dir) / "model_report.json").write_text(payload_json)
  90. total_models = sum(f["n_models"] for f in families)
  91. ctx.log.info(
  92. f"model_report: wrote model_report.pdf "
  93. f"({len(families)} family/families, {total_models} model(s))."
  94. )