model_report.py 3.8 KB

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