evalrx 0.1.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (339) hide show
  1. evalrx/__init__.py +139 -0
  2. evalrx/agent_assets/__init__.py +2 -0
  3. evalrx/agent_assets/skills/README.md +28 -0
  4. evalrx/agent_assets/skills/eval-chart-style/SKILL.md +172 -0
  5. evalrx/agent_assets/skills/evalrx-report-ui/SKILL.md +116 -0
  6. evalrx/agent_assets/skills/nature-figure/LICENSE +201 -0
  7. evalrx/agent_assets/skills/nature-figure/README.md +412 -0
  8. evalrx/agent_assets/skills/nature-figure/SKILL.md +60 -0
  9. evalrx/agent_assets/skills/nature-figure/manifest.yaml +59 -0
  10. evalrx/agent_assets/skills/nature-figure/references/api.md +436 -0
  11. evalrx/agent_assets/skills/nature-figure/references/backend-selection.md +100 -0
  12. evalrx/agent_assets/skills/nature-figure/references/chart-types.md +281 -0
  13. evalrx/agent_assets/skills/nature-figure/references/common-patterns.md +350 -0
  14. evalrx/agent_assets/skills/nature-figure/references/demos.md +65 -0
  15. evalrx/agent_assets/skills/nature-figure/references/design-theory.md +439 -0
  16. evalrx/agent_assets/skills/nature-figure/references/figure-contract.md +93 -0
  17. evalrx/agent_assets/skills/nature-figure/references/figure-legend-conventions.md +71 -0
  18. evalrx/agent_assets/skills/nature-figure/references/nature-2026-observations.md +112 -0
  19. evalrx/agent_assets/skills/nature-figure/references/qa-contract.md +119 -0
  20. evalrx/agent_assets/skills/nature-figure/references/r-template-index.md +66 -0
  21. evalrx/agent_assets/skills/nature-figure/references/r-workflow.md +161 -0
  22. evalrx/agent_assets/skills/nature-figure/references/tutorials.md +251 -0
  23. evalrx/agent_assets/skills/nature-figure/static/core/contract.md +29 -0
  24. evalrx/agent_assets/skills/nature-figure/static/core/stance.md +37 -0
  25. evalrx/agent_assets/skills/nature-figure/static/fragments/backend/python.md +37 -0
  26. evalrx/agent_assets/skills/nature-figure/static/fragments/backend/r.md +44 -0
  27. evalrx/agent_assets/skills/outcome-driver-analysis/SKILL.md +213 -0
  28. evalrx/agent_assets/skills/outcome-driver-analysis/assets/analysis_report_template.md +53 -0
  29. evalrx/agent_assets/skills/outcome-driver-analysis/references/model_selection.md +72 -0
  30. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/explanatory_var_eda.R +130 -0
  31. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/explanatory_var_eda.py +150 -0
  32. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/fit_outcome_model.R +181 -0
  33. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/fit_outcome_model.py +186 -0
  34. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/univariate_eda.R +149 -0
  35. evalrx/agent_assets/skills/outcome-driver-analysis/scripts/univariate_eda.py +177 -0
  36. evalrx/agent_assets/skills.py +27 -0
  37. evalrx/agent_runtime/__init__.py +78 -0
  38. evalrx/agent_runtime/_docker_runner.py +89 -0
  39. evalrx/agent_runtime/cli_runtime.py +103 -0
  40. evalrx/agent_runtime/cli_transcript.py +138 -0
  41. evalrx/agent_runtime/cli_types.py +68 -0
  42. evalrx/agent_runtime/codegen/__init__.py +5 -0
  43. evalrx/agent_runtime/codegen/runner.py +94 -0
  44. evalrx/agent_runtime/experiment_harness.py +117 -0
  45. evalrx/agent_runtime/factory.py +102 -0
  46. evalrx/agent_runtime/json_shape.py +44 -0
  47. evalrx/agent_runtime/judges/__init__.py +28 -0
  48. evalrx/agent_runtime/judges/agy.py +179 -0
  49. evalrx/agent_runtime/judges/autodetect.py +135 -0
  50. evalrx/agent_runtime/judges/claude.py +159 -0
  51. evalrx/agent_runtime/judges/codex.py +120 -0
  52. evalrx/agent_runtime/providers/__init__.py +21 -0
  53. evalrx/agent_runtime/providers/antigravity.py +31 -0
  54. evalrx/agent_runtime/providers/base.py +145 -0
  55. evalrx/agent_runtime/providers/claude_code.py +49 -0
  56. evalrx/agent_runtime/providers/codex.py +37 -0
  57. evalrx/agent_runtime/providers/gemini_cli.py +26 -0
  58. evalrx/agent_runtime/providers/kimi_cli.py +27 -0
  59. evalrx/agent_runtime/providers/opencode.py +27 -0
  60. evalrx/agent_runtime/providers/registry.py +58 -0
  61. evalrx/agent_runtime/sandbox.py +517 -0
  62. evalrx/agent_runtime/skill_audit.py +143 -0
  63. evalrx/agent_runtime/skills/__init__.py +19 -0
  64. evalrx/agent_runtime/skills/installer.py +68 -0
  65. evalrx/agent_runtime/skills/prompt_policy.py +86 -0
  66. evalrx/agent_runtime/skills/resolver.py +19 -0
  67. evalrx/analysis/__init__.py +132 -0
  68. evalrx/analysis/adjudicate.py +154 -0
  69. evalrx/analysis/analysis_module.py +361 -0
  70. evalrx/analysis/api.py +171 -0
  71. evalrx/analysis/case_studio.py +651 -0
  72. evalrx/analysis/cli.py +114 -0
  73. evalrx/analysis/dashboard.py +350 -0
  74. evalrx/analysis/eval_case_matrix.py +118 -0
  75. evalrx/analysis/eval_viz_theme.py +833 -0
  76. evalrx/analysis/explore_run.py +333 -0
  77. evalrx/analysis/explorer.py +1276 -0
  78. evalrx/analysis/failure_modes.py +607 -0
  79. evalrx/analysis/fused_pipeline.py +489 -0
  80. evalrx/analysis/holdout.py +300 -0
  81. evalrx/analysis/hypothesis_agent.py +230 -0
  82. evalrx/analysis/narration.py +177 -0
  83. evalrx/analysis/operationalize.py +442 -0
  84. evalrx/analysis/plain_language.py +42 -0
  85. evalrx/analysis/planner.py +283 -0
  86. evalrx/analysis/probe_search.py +203 -0
  87. evalrx/analysis/profile.py +268 -0
  88. evalrx/analysis/prompts/__init__.py +0 -0
  89. evalrx/analysis/prompts/explorer.py +417 -0
  90. evalrx/analysis/prompts/failure_modes.py +33 -0
  91. evalrx/analysis/prompts/holdout.py +27 -0
  92. evalrx/analysis/prompts/hypothesis_agent.py +78 -0
  93. evalrx/analysis/prompts/run_codebase.py +47 -0
  94. evalrx/analysis/prompts/stats_agent.py +72 -0
  95. evalrx/analysis/prompts/stats_tool_generator.py +43 -0
  96. evalrx/analysis/result_marker.py +47 -0
  97. evalrx/analysis/run_codebase.py +242 -0
  98. evalrx/analysis/run_view.py +205 -0
  99. evalrx/analysis/stage_views.py +93 -0
  100. evalrx/analysis/stats_agent.py +944 -0
  101. evalrx/analysis/stats_tool_agent.py +261 -0
  102. evalrx/analysis/stats_tool_generator.py +415 -0
  103. evalrx/analysis/stats_tools.py +1153 -0
  104. evalrx/analysis/trajectory_records.py +193 -0
  105. evalrx/analysis/workbench.py +431 -0
  106. evalrx/analyzers/__init__.py +42 -0
  107. evalrx/analyzers/agent/__init__.py +25 -0
  108. evalrx/analyzers/agent/counterfactual.py +84 -0
  109. evalrx/analyzers/agent/first_error_judge.py +96 -0
  110. evalrx/analyzers/agent/ignored_obs.py +81 -0
  111. evalrx/analyzers/agent/loop_detect.py +79 -0
  112. evalrx/analyzers/agent/reliability.py +165 -0
  113. evalrx/analyzers/agent/tool_shap.py +225 -0
  114. evalrx/analyzers/agent/trajectory_rubric.py +168 -0
  115. evalrx/analyzers/attention/__init__.py +19 -0
  116. evalrx/analyzers/attention/relative_attn.py +610 -0
  117. evalrx/analyzers/attention/rollout.py +73 -0
  118. evalrx/analyzers/attention/sink.py +56 -0
  119. evalrx/analyzers/attention/summary.py +190 -0
  120. evalrx/analyzers/attribution/__init__.py +6 -0
  121. evalrx/analyzers/attribution/generic_attn.py +31 -0
  122. evalrx/analyzers/attribution/gradcam.py +30 -0
  123. evalrx/analyzers/base.py +12 -0
  124. evalrx/analyzers/geometry/__init__.py +6 -0
  125. evalrx/analyzers/geometry/cka.py +70 -0
  126. evalrx/analyzers/geometry/linear_probe.py +157 -0
  127. evalrx/analyzers/hallucination/__init__.py +9 -0
  128. evalrx/analyzers/hallucination/chair.py +78 -0
  129. evalrx/analyzers/hallucination/opera.py +29 -0
  130. evalrx/analyzers/hallucination/pope.py +119 -0
  131. evalrx/analyzers/hallucination/selfcheck.py +155 -0
  132. evalrx/analyzers/hallucination/vcd.py +29 -0
  133. evalrx/analyzers/lens/__init__.py +7 -0
  134. evalrx/analyzers/lens/layer_contrast.py +133 -0
  135. evalrx/analyzers/lens/logit_lens.py +138 -0
  136. evalrx/analyzers/lens/tuned_lens.py +30 -0
  137. evalrx/analyzers/patching/__init__.py +5 -0
  138. evalrx/analyzers/patching/causal_trace.py +30 -0
  139. evalrx/analyzers/perturbation/__init__.py +23 -0
  140. evalrx/analyzers/perturbation/_shapley.py +54 -0
  141. evalrx/analyzers/perturbation/context_shap.py +174 -0
  142. evalrx/analyzers/perturbation/cot_faithfulness.py +239 -0
  143. evalrx/analyzers/perturbation/format_sensitivity.py +237 -0
  144. evalrx/analyzers/perturbation/mm_shap.py +146 -0
  145. evalrx/analyzers/perturbation/modality_ablation.py +196 -0
  146. evalrx/analyzers/perturbation/perturbation_battery.py +274 -0
  147. evalrx/analyzers/perturbation/prompt_contrast.py +265 -0
  148. evalrx/analyzers/perturbation/rise.py +94 -0
  149. evalrx/analyzers/perturbation/vl_shap.py +102 -0
  150. evalrx/analyzers/reasoning/__init__.py +33 -0
  151. evalrx/analyzers/reasoning/_text.py +328 -0
  152. evalrx/analyzers/reasoning/answer_extraction_audit.py +327 -0
  153. evalrx/analyzers/reasoning/arith_audit.py +226 -0
  154. evalrx/analyzers/reasoning/contamination.py +214 -0
  155. evalrx/analyzers/reasoning/knowledge_split.py +253 -0
  156. evalrx/analyzers/reasoning/self_repair.py +246 -0
  157. evalrx/analyzers/reasoning/step_rollout_value.py +216 -0
  158. evalrx/analyzers/reasoning/termination_audit.py +258 -0
  159. evalrx/analyzers/uncertainty/__init__.py +18 -0
  160. evalrx/analyzers/uncertainty/calibration.py +174 -0
  161. evalrx/analyzers/uncertainty/coverage_gap.py +199 -0
  162. evalrx/analyzers/uncertainty/entropy.py +90 -0
  163. evalrx/analyzers/uncertainty/logprob_entropy.py +69 -0
  164. evalrx/analyzers/uncertainty/self_consistency.py +204 -0
  165. evalrx/analyzers/uncertainty/verbalized_conf.py +64 -0
  166. evalrx/cli.py +411 -0
  167. evalrx/config.py +77 -0
  168. evalrx/contract/__init__.py +179 -0
  169. evalrx/contract/common.py +452 -0
  170. evalrx/contract/emit.py +948 -0
  171. evalrx/contract/export.py +237 -0
  172. evalrx/contract/m1.py +325 -0
  173. evalrx/contract/m2.py +317 -0
  174. evalrx/contract/m3.py +165 -0
  175. evalrx/contract/m4.py +130 -0
  176. evalrx/contract/m5.py +292 -0
  177. evalrx/contract/methodology.py +76 -0
  178. evalrx/contract/pre_m1.py +58 -0
  179. evalrx/contract/typescript.py +140 -0
  180. evalrx/core/__init__.py +85 -0
  181. evalrx/core/analyzer.py +174 -0
  182. evalrx/core/capability.py +54 -0
  183. evalrx/core/case.py +443 -0
  184. evalrx/core/experiment.py +106 -0
  185. evalrx/core/model.py +198 -0
  186. evalrx/core/pipeline.py +42 -0
  187. evalrx/core/registry.py +142 -0
  188. evalrx/core/result.py +64 -0
  189. evalrx/core/spec.py +173 -0
  190. evalrx/core/tokentype.py +165 -0
  191. evalrx/core/tool.py +92 -0
  192. evalrx/datasets/__init__.py +41 -0
  193. evalrx/datasets/base.py +68 -0
  194. evalrx/datasets/gui_os.py +52 -0
  195. evalrx/datasets/llm_qa.py +57 -0
  196. evalrx/datasets/pure_qa.py +12 -0
  197. evalrx/datasets/vlm_qa.py +695 -0
  198. evalrx/datasets/web_search_qa.py +52 -0
  199. evalrx/eval_agent/__init__.py +341 -0
  200. evalrx/eval_agent/_tools.py +81 -0
  201. evalrx/eval_agent/ab_runner.py +50 -0
  202. evalrx/eval_agent/agentic/__init__.py +43 -0
  203. evalrx/eval_agent/agentic/actions.py +216 -0
  204. evalrx/eval_agent/agentic/board.py +107 -0
  205. evalrx/eval_agent/agentic/loop.py +190 -0
  206. evalrx/eval_agent/agentic/tools.py +538 -0
  207. evalrx/eval_agent/checkpoint.py +57 -0
  208. evalrx/eval_agent/cli_agent.py +59 -0
  209. evalrx/eval_agent/cli_skills.py +5 -0
  210. evalrx/eval_agent/evolution.py +396 -0
  211. evalrx/eval_agent/git_manager.py +215 -0
  212. evalrx/eval_agent/hypothesis.py +172 -0
  213. evalrx/eval_agent/label_quarantine.py +209 -0
  214. evalrx/eval_agent/legacy.py +530 -0
  215. evalrx/eval_agent/log_schema.py +497 -0
  216. evalrx/eval_agent/loop.py +2159 -0
  217. evalrx/eval_agent/loop_reports.py +116 -0
  218. evalrx/eval_agent/model_instrumentation.py +282 -0
  219. evalrx/eval_agent/narration.py +193 -0
  220. evalrx/eval_agent/nl_runner.py +460 -0
  221. evalrx/eval_agent/orchestrator.py +61 -0
  222. evalrx/eval_agent/preregister.py +93 -0
  223. evalrx/eval_agent/prompts/__init__.py +1 -0
  224. evalrx/eval_agent/prompts/agentic.py +46 -0
  225. evalrx/eval_agent/prompts/case_discovery.py +25 -0
  226. evalrx/eval_agent/prompts/diagnosis.py +125 -0
  227. evalrx/eval_agent/prompts/experiment_writer.py +265 -0
  228. evalrx/eval_agent/prompts/explore_step.py +37 -0
  229. evalrx/eval_agent/prompts/fix_agent.py +257 -0
  230. evalrx/eval_agent/prompts/hypothesis_tester.py +15 -0
  231. evalrx/eval_agent/prompts/nl_runner.py +38 -0
  232. evalrx/eval_agent/prompts/probe_agent.py +25 -0
  233. evalrx/eval_agent/prompts/probe_candidate_generator.py +14 -0
  234. evalrx/eval_agent/prompts/probe_generator.py +35 -0
  235. evalrx/eval_agent/prompts/whitebox_probe_generator.py +38 -0
  236. evalrx/eval_agent/report.py +58 -0
  237. evalrx/eval_agent/run_context.py +354 -0
  238. evalrx/eval_agent/run_log.schema.json +1215 -0
  239. evalrx/eval_agent/run_logger_v2.py +1764 -0
  240. evalrx/eval_agent/run_metadata.py +208 -0
  241. evalrx/eval_agent/stages/__init__.py +56 -0
  242. evalrx/eval_agent/stages/case_discovery.py +293 -0
  243. evalrx/eval_agent/stages/diagnosis.py +1017 -0
  244. evalrx/eval_agent/stages/experiment_writer.py +1634 -0
  245. evalrx/eval_agent/stages/fix_agent.py +3916 -0
  246. evalrx/eval_agent/stages/fix_internals.py +499 -0
  247. evalrx/eval_agent/stages/fix_pipeline.py +725 -0
  248. evalrx/eval_agent/stages/fix_tiers.py +187 -0
  249. evalrx/eval_agent/stages/fix_tools.py +1034 -0
  250. evalrx/eval_agent/stages/hypothesis_tester.py +1014 -0
  251. evalrx/eval_agent/stages/probe.py +439 -0
  252. evalrx/eval_agent/stages/probe_agent.py +1079 -0
  253. evalrx/eval_agent/stages/probe_candidate_generator.py +128 -0
  254. evalrx/eval_agent/stages/probe_generator.py +326 -0
  255. evalrx/eval_agent/stages/probe_search_agent.py +106 -0
  256. evalrx/eval_agent/stages/protocol.py +112 -0
  257. evalrx/eval_agent/stages/repair_catalog.py +273 -0
  258. evalrx/eval_agent/stages/surgery.py +524 -0
  259. evalrx/eval_agent/stages/whitebox_probe_generator.py +351 -0
  260. evalrx/eval_agent/store.py +231 -0
  261. evalrx/logging_utils.py +112 -0
  262. evalrx/models/__init__.py +161 -0
  263. evalrx/models/_discover.py +101 -0
  264. evalrx/models/agent.py +380 -0
  265. evalrx/models/backends/__init__.py +58 -0
  266. evalrx/models/backends/api.py +169 -0
  267. evalrx/models/backends/base.py +57 -0
  268. evalrx/models/backends/gemini_compat.py +579 -0
  269. evalrx/models/backends/hf_local.py +2074 -0
  270. evalrx/models/backends/openai_compat.py +301 -0
  271. evalrx/models/backends/vllm_offline.py +116 -0
  272. evalrx/models/base.py +24 -0
  273. evalrx/models/blackbox/__init__.py +4 -0
  274. evalrx/models/blackbox/agent.py +31 -0
  275. evalrx/models/blackbox/base.py +29 -0
  276. evalrx/models/blackbox/gemini.py +279 -0
  277. evalrx/models/blackbox/llm_api.py +17 -0
  278. evalrx/models/blackbox/vlm_api.py +17 -0
  279. evalrx/models/compose.py +66 -0
  280. evalrx/models/inference.py +88 -0
  281. evalrx/models/paper_methods/__init__.py +8 -0
  282. evalrx/models/paper_methods/aad.py +53 -0
  283. evalrx/models/paper_methods/ifcd.py +204 -0
  284. evalrx/models/paper_methods/pai.py +164 -0
  285. evalrx/models/paper_methods/tcd.py +202 -0
  286. evalrx/models/paper_methods/vcd.py +45 -0
  287. evalrx/models/paper_methods/vicrop.py +137 -0
  288. evalrx/models/toolcodec.py +143 -0
  289. evalrx/models/tools/__init__.py +20 -0
  290. evalrx/models/tools/perception.py +300 -0
  291. evalrx/models/tools/visual.py +174 -0
  292. evalrx/models/whitebox/__init__.py +26 -0
  293. evalrx/models/whitebox/agent.py +31 -0
  294. evalrx/models/whitebox/base.py +24 -0
  295. evalrx/models/whitebox/qwen.py +61 -0
  296. evalrx/models/whitebox/qwen2_5_omni.py +29 -0
  297. evalrx/models/whitebox/qwen2_audio.py +25 -0
  298. evalrx/models/whitebox/qwen_omni.py +53 -0
  299. evalrx/models/whitebox/qwen_vl.py +62 -0
  300. evalrx/observability/__init__.py +21 -0
  301. evalrx/observability/envelope.py +122 -0
  302. evalrx/observability/outbox.py +111 -0
  303. evalrx/observability/tracer.py +882 -0
  304. evalrx/reporting/__init__.py +28 -0
  305. evalrx/reporting/case_study.py +947 -0
  306. evalrx/reporting/compiler.py +587 -0
  307. evalrx/reporting/dynamic.py +1882 -0
  308. evalrx/reporting/html_report.py +2225 -0
  309. evalrx/reporting/langfuse_exporter.py +38 -0
  310. evalrx/reporting/langfuse_source.py +155 -0
  311. evalrx/reporting/model.py +151 -0
  312. evalrx/reporting/run_events.py +184 -0
  313. evalrx/reporting/server.py +557 -0
  314. evalrx/reporting/stages.py +58 -0
  315. evalrx/reporting/static_export.py +142 -0
  316. evalrx/reporting/web_dist/index.html +146 -0
  317. evalrx/specs.py +727 -0
  318. evalrx/stats/__init__.py +47 -0
  319. evalrx/stats/api.py +192 -0
  320. evalrx/stats/bootstrap.py +86 -0
  321. evalrx/stats/ebh.py +27 -0
  322. evalrx/stats/evalue.py +98 -0
  323. evalrx/stats/friedman.py +138 -0
  324. evalrx/stats/mcnemar.py +40 -0
  325. evalrx/stats/multiplicity.py +159 -0
  326. evalrx/stats/subset_sampling.py +55 -0
  327. evalrx/term_links.py +43 -0
  328. evalrx/viz/__init__.py +7 -0
  329. evalrx/viz/labels.py +77 -0
  330. evalrx/viz/prompts.py +39 -0
  331. evalrx/viz/renderer.py +590 -0
  332. evalrx/viz/schema.py +36 -0
  333. evalrx/viz/style.py +134 -0
  334. evalrx-0.1.2.dist-info/METADATA +532 -0
  335. evalrx-0.1.2.dist-info/RECORD +339 -0
  336. evalrx-0.1.2.dist-info/WHEEL +5 -0
  337. evalrx-0.1.2.dist-info/entry_points.txt +3 -0
  338. evalrx-0.1.2.dist-info/licenses/LICENSE +121 -0
  339. evalrx-0.1.2.dist-info/top_level.txt +1 -0
@@ -0,0 +1,1153 @@
1
+ """M2 — statistical tool catalog: thin wrappers around :mod:`evalrx.stats`.
2
+
3
+ This module turns the rigorous-but-low-level :mod:`evalrx.stats` building
4
+ blocks (McNemar + e-value, clustered bootstrap, Friedman/Nemenyi, e-BH) into a
5
+ small **catalog of named tools** that
6
+ :class:`~evalrx.analysis.stats_agent.StatsAnalysisAgent` can
7
+ *select* and call with a config dict — the "select stats tools" half of the M2
8
+ plan (Plan A, 2026-06-05).
9
+
10
+ The flow is:
11
+
12
+ 1. :func:`build_stats_input` normalises ``{analyzer: Result}`` + a labeled
13
+ ``CaseBatch`` into a single :class:`StatsInput` (per-case signals, labels,
14
+ scalar metrics, optional strategy groups).
15
+ 2. The agent picks tool names from :data:`STATS_TOOL_CATALOG` (LLM-guided) or
16
+ falls back to :func:`default_plan` (deterministic).
17
+ 3. Each ``(tool, config)`` runs via :func:`run_stats_tool`, returning a uniform
18
+ :class:`StatsToolResult` (effect, CI, e-value, reject, underpowered).
19
+ 4. :func:`fdr_correct` applies e-BH across all tools that produced an e-value so
20
+ multiple-metric testing is FDR-controlled, not naive.
21
+ 5. :func:`plot_effects` draws an optional forest plot of effect ± CI.
22
+
23
+ No tool ever returns a bare p-value: every verdict carries an effect size and a
24
+ corrected reject decision, inherited from :func:`evalrx.stats.compare`.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import logging
30
+ import math
31
+ import os
32
+ from dataclasses import dataclass, field
33
+ from typing import TYPE_CHECKING, Any, Callable
34
+
35
+ import numpy as np
36
+
37
+ from evalrx.core.case import Label
38
+ from evalrx.stats import (
39
+ compare,
40
+ compare_multiple,
41
+ e_value_test,
42
+ kendall_tau,
43
+ )
44
+ from evalrx.stats.multiplicity import correct_results
45
+
46
+ if TYPE_CHECKING:
47
+ from evalrx.core.case import CaseBatch
48
+ from evalrx.core.result import Result
49
+
50
+ logger = logging.getLogger(__name__)
51
+
52
+ # Keys in a per_case finding entry that identify the case, not a signal.
53
+ _ID_KEYS = ("sample_id", "case_id", "id")
54
+
55
+
56
+ # ---------------------------------------------------------------------------
57
+ # Normalised input + uniform result
58
+ # ---------------------------------------------------------------------------
59
+
60
+ @dataclass
61
+ class StatsInput:
62
+ """Normalised view of M1 results + labels, ready for statistical tests.
63
+
64
+ Attributes:
65
+ labels: ``{case_id -> is_fail}`` for PASS/FAIL cases (UNKNOWN dropped).
66
+ per_case: ``{"analyzer.metric" -> {case_id -> value}}`` — per-case
67
+ numeric/boolean signals harvested from ``findings["per_case"]``.
68
+ scalars: ``{"analyzer.metric" -> value}`` — aggregate numeric findings.
69
+ groups: Optional ``{strategy -> {case_id -> success}}`` for
70
+ paired/omnibus strategy comparisons (from ``findings["by_strategy"]``).
71
+ """
72
+
73
+ labels: dict[str, bool] = field(default_factory=dict)
74
+ per_case: dict[str, dict[str, float]] = field(default_factory=dict)
75
+ scalars: dict[str, float] = field(default_factory=dict)
76
+ groups: dict[str, dict[str, float]] | None = None
77
+ # Per-case signals that near-perfectly RECONSTRUCT the FAIL label (a probe
78
+ # output equal to the label, a label-recomputing recipe, …). Moved here by
79
+ # :func:`isolate_label_leaks` so they never enter the tested family / e-BH
80
+ # multiplicity / candidate charts / hypothesis seeding — but are KEPT as a
81
+ # pipeline self-check (the plumbing audit). ``{name -> {case_id -> value}}``.
82
+ sanity: dict[str, dict[str, float]] = field(default_factory=dict)
83
+ # Per-case VECTOR signals (e.g. a full attention map per case), harvested from
84
+ # ``Result.artifacts["per_case_maps"]``. Consumed by the tensor-level
85
+ # ``attention_decoding`` omnibus, NOT the scalar tools.
86
+ # ``{name -> {case_id -> np.ndarray}}``.
87
+ per_case_vectors: dict[str, dict[str, "Any"]] = field(default_factory=dict)
88
+
89
+ @classmethod
90
+ def from_results(
91
+ cls,
92
+ results: "dict[str, Result]",
93
+ data: "CaseBatch | None" = None,
94
+ ) -> "StatsInput":
95
+ """Build statistical input from EvalRX analyzer results."""
96
+ return build_stats_input(results, data)
97
+
98
+ @classmethod
99
+ def from_records(
100
+ cls,
101
+ records: "Any",
102
+ *,
103
+ id_col: str = "case_id",
104
+ label_col: str = "label",
105
+ signal_cols: "list[str] | tuple[str, ...] | None" = None,
106
+ scalar_cols: "list[str] | tuple[str, ...] | None" = None,
107
+ signal_prefix: str = "",
108
+ ) -> "StatsInput":
109
+ """Build statistical input from plain row dictionaries.
110
+
111
+ ``label_col`` accepts booleans, ``Label`` values, common strings
112
+ (``"pass"``, ``"fail"``, ``"success"``, ``"error"``), or 0/1 values
113
+ where 1 means FAIL. Signal and scalar columns must be numeric/bool.
114
+ """
115
+ return build_stats_input_from_records(
116
+ records,
117
+ id_col=id_col,
118
+ label_col=label_col,
119
+ signal_cols=signal_cols,
120
+ scalar_cols=scalar_cols,
121
+ signal_prefix=signal_prefix,
122
+ )
123
+
124
+
125
+ @dataclass
126
+ class StatsToolResult:
127
+ """Uniform output of one statistical tool.
128
+
129
+ Attributes:
130
+ tool: Catalog name of the tool that produced this.
131
+ config: The (resolved) config the tool ran with.
132
+ ok: ``True`` when the test ran; ``False`` when skipped/failed.
133
+ summary: Human-readable one-liner.
134
+ effect: Effect size (e.g. fail-rate difference, τ) when applicable.
135
+ ci: Confidence interval on the effect, when applicable.
136
+ reject: Corrected reject decision (e-value / CI), never a bare p.
137
+ e_value: Anytime-valid e-value, when the test produces one.
138
+ p_value: Raw p (diagnostic only — never the decision basis).
139
+ underpowered: Inconclusive *and* CI too wide to rule out a real effect.
140
+ details: Tool-specific extras (group sizes, ranks, contingency …).
141
+ figure_path: Path to a per-tool figure, if any.
142
+ error: Reason string when ``ok`` is ``False``.
143
+ """
144
+
145
+ tool: str
146
+ config: dict[str, Any] = field(default_factory=dict)
147
+ ok: bool = True
148
+ summary: str = ""
149
+ effect: float | None = None
150
+ ci: tuple[float, float] | None = None
151
+ reject: bool | None = None
152
+ e_value: float | None = None
153
+ p_value: float | None = None
154
+ underpowered: bool = False
155
+ details: dict[str, Any] = field(default_factory=dict)
156
+ figure_path: str | None = None
157
+ error: str | None = None
158
+ # Generic-M2 metadata. Existing callers can ignore these fields; they let
159
+ # planners/controllers identify one result inside a larger tested family.
160
+ analysis_key: str | None = None
161
+ correction_family: str | None = "auto" # "e_bh" | "bh" | "auto" | None
162
+ correction_method: str | None = None
163
+ fdr_corrected: bool = False
164
+ raw_reject: bool | None = None
165
+
166
+ def to_dict(self) -> dict[str, Any]:
167
+ return {
168
+ "tool": self.tool,
169
+ "config": self.config,
170
+ "ok": self.ok,
171
+ "summary": self.summary,
172
+ "effect": self.effect,
173
+ "ci": list(self.ci) if self.ci is not None else None,
174
+ "reject": self.reject,
175
+ "e_value": self.e_value,
176
+ "p_value": self.p_value,
177
+ "underpowered": self.underpowered,
178
+ "details": self.details,
179
+ "figure_path": self.figure_path,
180
+ "error": self.error,
181
+ "analysis_key": self.analysis_key,
182
+ "correction_family": self.correction_family,
183
+ "correction_method": self.correction_method,
184
+ "fdr_corrected": self.fdr_corrected,
185
+ "raw_reject": self.raw_reject,
186
+ }
187
+
188
+
189
+ EvidenceResult = StatsToolResult
190
+
191
+
192
+ # ---------------------------------------------------------------------------
193
+ # Input construction
194
+ # ---------------------------------------------------------------------------
195
+
196
+ def _entry_id(entry: dict) -> str:
197
+ for k in _ID_KEYS:
198
+ v = entry.get(k)
199
+ if v:
200
+ return str(v)
201
+ return ""
202
+
203
+
204
+ def build_stats_input(
205
+ results: "dict[str, Result]",
206
+ data: "CaseBatch | None" = None,
207
+ ) -> StatsInput:
208
+ """Normalise ``{analyzer: Result}`` + optional labeled *data* into a :class:`StatsInput`."""
209
+ labels: dict[str, bool] = {}
210
+ if data is not None:
211
+ for c in data:
212
+ lab = getattr(c, "label", None)
213
+ if lab is None or lab == Label.UNKNOWN:
214
+ continue
215
+ is_fail = lab == Label.FAIL
216
+ labels[c.id] = is_fail
217
+ traj = getattr(c, "trajectory", None)
218
+ sid = getattr(traj, "sample_id", "") if traj is not None else ""
219
+ if sid:
220
+ labels[sid] = is_fail
221
+
222
+ per_case: dict[str, dict[str, float]] = {}
223
+ scalars: dict[str, float] = {}
224
+ groups: dict[str, dict[str, float]] = {}
225
+ per_case_vectors: dict[str, dict[str, Any]] = {}
226
+
227
+ for aname, res in results.items():
228
+ findings = res.findings or {}
229
+ for k, v in findings.items():
230
+ if isinstance(v, (int, float, bool)):
231
+ scalars[f"{aname}.{k}"] = float(v)
232
+ for entry in findings.get("per_case", []) or []:
233
+ if not isinstance(entry, dict):
234
+ continue
235
+ cid = _entry_id(entry)
236
+ if not cid:
237
+ continue
238
+ for k, v in entry.items():
239
+ if k in _ID_KEYS:
240
+ continue
241
+ if isinstance(v, (int, float, bool)):
242
+ per_case.setdefault(f"{aname}.{k}", {})[cid] = float(v)
243
+ by_strategy = findings.get("by_strategy")
244
+ if isinstance(by_strategy, dict):
245
+ for sname, vec in by_strategy.items():
246
+ slot = groups.setdefault(str(sname), {})
247
+ if isinstance(vec, dict):
248
+ for cid, val in vec.items():
249
+ if isinstance(val, (int, float, bool)):
250
+ slot[str(cid)] = float(bool(val))
251
+
252
+ # Per-case VECTOR signals (full attention maps) for the tensor-level
253
+ # omnibus — kept in artifacts (heavy), so read them off the Result here.
254
+ maps = (getattr(res, "artifacts", None) or {}).get("per_case_maps")
255
+ if isinstance(maps, dict) and maps:
256
+ col = {str(cid): m for cid, m in maps.items() if m is not None}
257
+ if col:
258
+ per_case_vectors[f"{aname}.map"] = col
259
+
260
+ out = StatsInput(
261
+ labels=labels,
262
+ per_case=per_case,
263
+ scalars=scalars,
264
+ groups=groups or None,
265
+ per_case_vectors=per_case_vectors,
266
+ )
267
+ # Route label-reconstructing signals to the sanity lane so they never enter
268
+ # the tested family / e-BH multiplicity / candidate charts.
269
+ isolate_label_leaks(out)
270
+ return out
271
+
272
+
273
+ def build_stats_input_from_records(
274
+ records: "Any",
275
+ *,
276
+ id_col: str = "case_id",
277
+ label_col: str = "label",
278
+ signal_cols: "list[str] | tuple[str, ...] | None" = None,
279
+ scalar_cols: "list[str] | tuple[str, ...] | None" = None,
280
+ signal_prefix: str = "",
281
+ ) -> StatsInput:
282
+ """Normalise plain records into :class:`StatsInput`.
283
+
284
+ This is the standalone on-ramp for users who have a table of cases and
285
+ signals rather than EvalRX ``Result`` objects.
286
+ """
287
+ labels: dict[str, bool] = {}
288
+ per_case: dict[str, dict[str, float]] = {}
289
+ scalars: dict[str, float] = {}
290
+
291
+ rows = list(records or [])
292
+ if signal_cols is None and rows:
293
+ excluded = {id_col, label_col, *(scalar_cols or ())}
294
+ signal_cols = [
295
+ str(k) for k, v in _row_items(rows[0])
296
+ if k not in excluded and isinstance(v, (int, float, bool))
297
+ ]
298
+ signal_cols = tuple(signal_cols or ())
299
+ scalar_cols = tuple(scalar_cols or ())
300
+
301
+ for i, row in enumerate(rows):
302
+ cid = _row_get(row, id_col, None)
303
+ if cid in (None, ""):
304
+ cid = str(i)
305
+ cid = str(cid)
306
+
307
+ label = _parse_label(_row_get(row, label_col, None))
308
+ if label is not None:
309
+ labels[cid] = label
310
+
311
+ for col in signal_cols:
312
+ val = _row_get(row, col, None)
313
+ if isinstance(val, (int, float, bool)):
314
+ key = f"{signal_prefix}.{col}" if signal_prefix else str(col)
315
+ per_case.setdefault(key, {})[cid] = float(val)
316
+
317
+ for col in scalar_cols:
318
+ val = _row_get(row, col, None)
319
+ if isinstance(val, (int, float, bool)):
320
+ scalars[str(col)] = float(val)
321
+
322
+ out = StatsInput(labels=labels, per_case=per_case, scalars=scalars)
323
+ isolate_label_leaks(out)
324
+ return out
325
+
326
+
327
+ def _row_get(row: Any, key: str, default: Any = None) -> Any:
328
+ if isinstance(row, dict):
329
+ return row.get(key, default)
330
+ return getattr(row, key, default)
331
+
332
+
333
+ def _row_items(row: Any) -> list[tuple[str, Any]]:
334
+ if isinstance(row, dict):
335
+ return list(row.items())
336
+ if hasattr(row, "_asdict"):
337
+ return list(row._asdict().items())
338
+ if hasattr(row, "__dict__"):
339
+ return list(vars(row).items())
340
+ return []
341
+
342
+
343
+ def _parse_label(value: Any) -> bool | None:
344
+ if value is None or value == Label.UNKNOWN:
345
+ return None
346
+ if value == Label.FAIL:
347
+ return True
348
+ if value == Label.PASS:
349
+ return False
350
+ if isinstance(value, bool):
351
+ return value
352
+ if isinstance(value, (int, float)):
353
+ return bool(value)
354
+ text = str(value).strip().lower()
355
+ if text in {"fail", "failed", "failure", "false", "incorrect", "error", "bad", "1"}:
356
+ return True
357
+ if text in {"pass", "passed", "success", "true", "correct", "ok", "0"}:
358
+ return False
359
+ return None
360
+
361
+
362
+ def _is_binary(values: "Any") -> bool:
363
+ return all(float(v) in (0.0, 1.0) for v in values)
364
+
365
+
366
+ def _binarize(
367
+ signal_map: dict[str, float],
368
+ mode: str = "median",
369
+ threshold: float | None = None,
370
+ ) -> dict[str, bool]:
371
+ """Binarise a continuous per-case signal (median split by default)."""
372
+ vals = list(signal_map.values())
373
+ if not vals or _is_binary(vals):
374
+ return {cid: bool(v) for cid, v in signal_map.items()}
375
+ if threshold is None:
376
+ if mode == "median":
377
+ s = sorted(vals)
378
+ n = len(s)
379
+ threshold = s[n // 2] if n % 2 else (s[n // 2 - 1] + s[n // 2]) / 2.0
380
+ else:
381
+ threshold = 0.0
382
+ return {cid: (v > threshold) for cid, v in signal_map.items()}
383
+
384
+
385
+ # ---------------------------------------------------------------------------
386
+ # Label-leak detection (the deferred "leak-1" check from operationalize.py): a
387
+ # per-case signal that RECONSTRUCTS the FAIL label carries no diagnostic info —
388
+ # it is the label in disguise (e.g. a probe whose output equals the failure
389
+ # definition). Detect such columns statistically and route them to a separate
390
+ # "sanity" lane instead of testing/charting them as discriminators.
391
+ # ---------------------------------------------------------------------------
392
+
393
+ # A leak is the label *in disguise* — NOT merely a strong predictor. The
394
+ # signature is a BINARY flag that ~equals the FAIL label (a recomputed outcome,
395
+ # e.g. a probe that re-derives "is this a false detection"). A CONTINUOUS feature
396
+ # that perfectly separates the classes (e.g. object size) is legitimate discovery,
397
+ # the very thing we want to find — so separation alone never flags it. Recipe-level
398
+ # label references are caught earlier by compile_recipe's G4 guard.
399
+ _LEAK_MIN_N = 10
400
+ # A binary signal matching the FAIL label at ≥0.95 is a recomputed outcome, not a
401
+ # mechanism — genuine binary mechanism signals are noisy, and a probe re-deriving
402
+ # the answer lands near 1.0 (minus a little label drift). 0.95 catches the latter
403
+ # robustly while leaving any merely-strong (≤0.9) binary feature in the family.
404
+ _LEAK_BINARY_ACC = 0.95
405
+
406
+
407
+ def _auc(scores: list[float], labels: list[int]) -> float:
408
+ """ROC-AUC of *scores* vs binary *labels* (rank-based, ties averaged)."""
409
+ n_pos = sum(labels)
410
+ n_neg = len(labels) - n_pos
411
+ if n_pos == 0 or n_neg == 0:
412
+ return 0.5
413
+ order = sorted(range(len(scores)), key=lambda i: scores[i])
414
+ ranks = [0.0] * len(scores)
415
+ i = 0
416
+ while i < len(order):
417
+ j = i
418
+ while j + 1 < len(order) and scores[order[j + 1]] == scores[order[i]]:
419
+ j += 1
420
+ avg = (i + j) / 2.0 + 1.0 # 1-based average rank across the tie block
421
+ for k in range(i, j + 1):
422
+ ranks[order[k]] = avg
423
+ i = j + 1
424
+ rank_pos = sum(ranks[idx] for idx, y in enumerate(labels) if y)
425
+ u = rank_pos - n_pos * (n_pos + 1) / 2.0
426
+ return u / (n_pos * n_neg)
427
+
428
+
429
+ def label_leak_score(sigmap: dict[str, float], labels: dict[str, bool]) -> dict[str, Any]:
430
+ """Score whether a per-case signal IS the FAIL label in disguise.
431
+
432
+ Returns ``{n, leak, score, kind, reason}``. Only BINARY signals can be flagged
433
+ (``score`` = best-split accuracy vs the label); a binary flag that matches the
434
+ label to ``_LEAK_BINARY_ACC`` with ≥ ``_LEAK_MIN_N`` cases is a recomputed
435
+ outcome, not a discriminator. CONTINUOUS signals are NEVER flagged — perfect
436
+ separation by a real feature is the discovery we want, not leakage (their AUC
437
+ margin is still reported as ``score`` for transparency).
438
+ """
439
+ vals = list(sigmap.values())
440
+ binary = bool(vals) and _is_binary(vals)
441
+ # Align to labeled cases. Sparse binary flags: a labeled case missing from the
442
+ # map means the signal is ABSENT (mirrors _split_signal_groups); continuous: skip.
443
+ xs: list[float] = []
444
+ ys: list[int] = []
445
+ for cid, is_fail in labels.items():
446
+ if cid in sigmap:
447
+ xs.append(float(sigmap[cid]))
448
+ elif binary:
449
+ xs.append(0.0)
450
+ else:
451
+ continue
452
+ ys.append(int(is_fail))
453
+ n = len(xs)
454
+ kind = "binary" if binary else "continuous"
455
+ if n < _LEAK_MIN_N or not any(ys) or all(ys):
456
+ return {"n": n, "leak": False, "score": 0.0, "kind": kind, "reason": ""}
457
+ if not binary:
458
+ # Report rank separation but never flag it — a perfectly separating
459
+ # continuous feature is a finding, not a leak.
460
+ margin = abs(2.0 * _auc(xs, ys) - 1.0)
461
+ return {"n": n, "leak": False, "score": round(margin, 4),
462
+ "kind": "continuous", "reason": ""}
463
+ agree = sum(1 for x, y in zip(xs, ys) if int(x > 0.5) == y) / n
464
+ acc = max(agree, 1.0 - agree) # the signal may track FAIL or track PASS
465
+ leak = acc >= _LEAK_BINARY_ACC
466
+ return {"n": n, "leak": leak, "score": round(acc, 4), "kind": "binary",
467
+ "reason": (f"binary signal reconstructs the FAIL label "
468
+ f"(best-split accuracy {acc:.3f})") if leak else ""}
469
+
470
+
471
+ #: Per-case signals that are the OUTCOME re-graded, not a mechanism: "is the
472
+ #: model's (baseline or strategy) answer correct" by an analyzer's own matcher.
473
+ #: They agree with the official label at 80-95% — below the leak threshold
474
+ #: (the matchers differ), above anything a real mechanism signal reaches — so
475
+ #: left in the family they are the BH survivors every time (qwen3.5-2b/
476
+ #: minervamath: all 6 survivors of 41 tests were these, effect -0.63 each) and
477
+ #: M4 then "refutes" or "supports" hypotheses on a tautology. The mechanism
478
+ #: content of these analyzers lives in their DERIVED flags (extraction_suspect,
479
+ #: label_disagrees, coverage_gap, majority_share, changed_answer, …), which
480
+ #: stay; strategy outcomes are compared PAIRED through ``groups``.
481
+ #: Matched on the metric name (the part after the analyzer prefix).
482
+ OUTCOME_REGRADE_METRICS: frozenset = frozenset({
483
+ # the baseline answer re-graded
484
+ "gold_in_output", "gold_in_answer_region", "strict_match", "answer_correct",
485
+ "baseline_correct", "final_correct", "is_correct",
486
+ # an intervention arm's answer graded (same item, same model: the
487
+ # association with the label is the baseline's; the arm's VALUE is the
488
+ # analyzer's gain scalar / derived flag, or a paired test)
489
+ "revised_correct", "decomposed_correct", "own_facts_correct",
490
+ "open_book_correct", "direct_correct", "reask_correct", "verify_correct",
491
+ "continuation_correct", "told_correct",
492
+ "majority_correct", "any_correct", "pass_at_k",
493
+ # the COUNT of correct samples among k (coverage_verification_gap): a
494
+ # 0..k integer, so label_leak_score (binary-only) never flags it, and
495
+ # under degenerate sampling it is exactly {0, k} = the label. It was the
496
+ # lone BH survivor on spatial457/qwen2.5-vl (2026-08-20) and named as
497
+ # M4 evidence. n_unique (sample diversity) stays: that is a mechanism.
498
+ "n_correct",
499
+ })
500
+
501
+
502
+ def _is_outcome_regrade(name: str) -> bool:
503
+ metric = str(name).rsplit(".", 1)[-1].lower()
504
+ return metric in OUTCOME_REGRADE_METRICS
505
+
506
+
507
+ def isolate_label_leaks(inp: StatsInput, *, denylist: "tuple[str, ...]" = ()) -> dict[str, str]:
508
+ """Move label-reconstructing per-case columns from ``per_case`` to ``sanity``.
509
+
510
+ Idempotent. A column is isolated when :func:`label_leak_score` flags it (a
511
+ near-perfect label stand-in), when it is an outcome re-grade
512
+ (:data:`OUTCOME_REGRADE_METRICS`), or its name contains a *denylist*
513
+ substring. Returns ``{name -> reason}`` for the moved columns so callers can
514
+ audit them. Leak-free columns are untouched, so the tested family holds only
515
+ genuine candidate discriminators — and the explorer (fed ``per_case``) won't
516
+ chart the isolated ones either.
517
+ """
518
+ moved: dict[str, str] = {}
519
+ for name in list(inp.per_case):
520
+ reason = ""
521
+ if denylist and any(d in name for d in denylist):
522
+ reason = "name matches leak denylist"
523
+ elif _is_outcome_regrade(name):
524
+ reason = "outcome re-grade (the answer's correctness, not a mechanism signal)"
525
+ else:
526
+ sc = label_leak_score(inp.per_case[name], inp.labels)
527
+ if sc["leak"]:
528
+ reason = sc["reason"]
529
+ if reason:
530
+ inp.sanity[name] = inp.per_case.pop(name)
531
+ moved[name] = reason
532
+ if moved:
533
+ logger.info("isolated %d label-reconstructing signal(s) to the sanity lane: %s",
534
+ len(moved), ", ".join(sorted(moved)))
535
+ return moved
536
+
537
+
538
+ def describe_data(inp: StatsInput) -> dict[str, Any]:
539
+ """Compact, LLM-friendly summary of what statistical tests are feasible."""
540
+ n_fail = sum(1 for v in inp.labels.values() if v)
541
+ n_labeled = len(inp.labels)
542
+ continuous = [k for k, m in inp.per_case.items() if not _is_binary(m.values())]
543
+ return {
544
+ "n_labeled": n_labeled,
545
+ "n_fail": n_fail,
546
+ "n_pass": n_labeled - n_fail,
547
+ "per_case_signals": list(inp.per_case),
548
+ "continuous_signals": continuous,
549
+ "scalar_metrics": list(inp.scalars),
550
+ "n_strategy_groups": len(inp.groups) if inp.groups else 0,
551
+ # Label-reconstructing signals held out of the tested family (audit only).
552
+ "sanity_signals": list(inp.sanity),
553
+ }
554
+
555
+
556
+ def has_testable_data(inp: StatsInput) -> bool:
557
+ """True when at least one tool could run (labels or strategy groups present)."""
558
+ return bool(inp.labels) or bool(inp.groups)
559
+
560
+
561
+ # ---------------------------------------------------------------------------
562
+ # Tool implementations
563
+ # ---------------------------------------------------------------------------
564
+
565
+ def _split_signal_groups(
566
+ inp: StatsInput,
567
+ signal_key: str,
568
+ binarize: str = "median",
569
+ threshold: float | None = None,
570
+ ) -> tuple[list[int], list[int]] | None:
571
+ """Split labeled cases into (signal-present, signal-absent) fail indicators.
572
+
573
+ A sparse binary flag only lists the cases where it fired, so a labeled case
574
+ missing from the signal map means the signal was *absent* (control group).
575
+ For a continuous signal we cannot assume a value, so missing cases are
576
+ skipped rather than defaulted.
577
+ """
578
+ sigmap = inp.per_case.get(signal_key)
579
+ if not sigmap:
580
+ return None
581
+ binar = _binarize(sigmap, binarize, threshold)
582
+ treat_missing_as_absent = _is_binary(sigmap.values())
583
+ signal_fail: list[int] = []
584
+ control_fail: list[int] = []
585
+ for cid, is_fail in inp.labels.items():
586
+ if cid in binar:
587
+ present = binar[cid]
588
+ elif treat_missing_as_absent:
589
+ present = False
590
+ else:
591
+ continue
592
+ (signal_fail if present else control_fail).append(int(is_fail))
593
+ return signal_fail, control_fail
594
+
595
+
596
+ def _mean(xs: list[int]) -> float:
597
+ return sum(xs) / len(xs) if xs else 0.0
598
+
599
+
600
+ def _two_group_permutation_p(control_fail: list[int], signal_fail: list[int]) -> float:
601
+ """Exact two-sided permutation p-value for a binary outcome split.
602
+
603
+ Conditions on the group sizes and total number of failures. This gives the
604
+ marginal signal test a p-value family so generalized M2 can control BH across
605
+ many candidate signals.
606
+ """
607
+ n_control = len(control_fail)
608
+ n_signal = len(signal_fail)
609
+ n = n_control + n_signal
610
+ if not n_control or not n_signal:
611
+ return 1.0
612
+ total_fail = sum(control_fail) + sum(signal_fail)
613
+ obs = abs(_mean(signal_fail) - _mean(control_fail))
614
+ denom = math.comb(n, n_signal)
615
+ if denom == 0:
616
+ return 1.0
617
+ lo = max(0, n_signal - (n - total_fail))
618
+ hi = min(n_signal, total_fail)
619
+ prob = 0.0
620
+ eps = 1e-12
621
+ for k in range(lo, hi + 1):
622
+ signal_rate = k / n_signal
623
+ control_rate = (total_fail - k) / n_control
624
+ if abs(signal_rate - control_rate) + eps >= obs:
625
+ prob += math.comb(total_fail, k) * math.comb(n - total_fail, n_signal - k) / denom
626
+ return min(1.0, max(0.0, prob))
627
+
628
+
629
+ def _tool_signal_label_assoc(inp: StatsInput, config: dict) -> StatsToolResult:
630
+ """Unpaired fail-rate difference between cases with/without a per-case signal."""
631
+ key = config.get("signal") or next(iter(inp.per_case), None)
632
+ cfg = {**config, "signal": key}
633
+ if not key or key not in inp.per_case:
634
+ return StatsToolResult(
635
+ tool="signal_label_assoc", config=cfg, ok=False,
636
+ error="no per-case signal available", summary="signal_label_assoc: no signal",
637
+ )
638
+ split = _split_signal_groups(
639
+ inp, key, config.get("binarize", "median"), config.get("threshold")
640
+ )
641
+ assert split is not None
642
+ signal_fail, control_fail = split
643
+ if not signal_fail or not control_fail:
644
+ return StatsToolResult(
645
+ tool="signal_label_assoc", config=cfg, ok=False,
646
+ error="one group empty (need both signal-present and signal-absent cases)",
647
+ summary="signal_label_assoc: degenerate split",
648
+ details={"n_signal": len(signal_fail), "n_control": len(control_fail)},
649
+ )
650
+ sr = compare(
651
+ control_fail, signal_fail, paired=False,
652
+ alpha=config.get("alpha", 0.05),
653
+ min_effect=config.get("min_effect", 0.0),
654
+ n_boot=config.get("n_boot", 2000),
655
+ )
656
+ p_value = _two_group_permutation_p(control_fail, signal_fail)
657
+ return StatsToolResult(
658
+ tool="signal_label_assoc", config=cfg, ok=True,
659
+ effect=sr.effect, ci=sr.ci, reject=sr.reject, p_value=p_value,
660
+ underpowered=sr.underpowered,
661
+ summary=f"signal '{key}' vs FAIL: {sr.summary()}",
662
+ analysis_key=f"signal_label_assoc:{key}",
663
+ correction_family="bh",
664
+ raw_reject=sr.reject,
665
+ details={
666
+ "n_signal": len(signal_fail), "n_control": len(control_fail),
667
+ "fail_rate_signal": round(_mean(signal_fail), 4),
668
+ "fail_rate_control": round(_mean(control_fail), 4),
669
+ "permutation_p": round(p_value, 6),
670
+ **sr.details,
671
+ },
672
+ )
673
+
674
+
675
+ def _aligned_groups(inp: StatsInput, names: list[str]) -> tuple[list[str], list[list[float]]]:
676
+ sets = [set(inp.groups[n]) for n in names] # type: ignore[index]
677
+ common = sorted(set.intersection(*sets)) if sets else []
678
+ vecs = [[inp.groups[n][cid] for cid in common] for n in names] # type: ignore[index]
679
+ return common, vecs
680
+
681
+
682
+ def _tool_bootstrap_diff(inp: StatsInput, config: dict) -> StatsToolResult:
683
+ """Unpaired fail-rate difference between two strategy groups (bootstrap CI)."""
684
+ if not inp.groups or len(inp.groups) < 2:
685
+ return StatsToolResult(
686
+ tool="bootstrap_diff", config=config, ok=False,
687
+ error="need >=2 strategy groups", summary="bootstrap_diff: <2 groups",
688
+ )
689
+ names = config.get("strategies") or list(inp.groups)[:2]
690
+ names = list(names)[:2]
691
+ common, vecs = _aligned_groups(inp, names)
692
+ if not common:
693
+ return StatsToolResult(
694
+ tool="bootstrap_diff", config={**config, "strategies": names}, ok=False,
695
+ error="no shared cases between groups", summary="bootstrap_diff: no overlap",
696
+ )
697
+ a, b = vecs
698
+ sr = compare(
699
+ a, b, paired=False,
700
+ alpha=config.get("alpha", 0.05), min_effect=config.get("min_effect", 0.0),
701
+ n_boot=config.get("n_boot", 2000),
702
+ )
703
+ return StatsToolResult(
704
+ tool="bootstrap_diff", config={**config, "strategies": names}, ok=True,
705
+ effect=sr.effect, ci=sr.ci, reject=sr.reject, underpowered=sr.underpowered,
706
+ summary=f"{names[1]} vs {names[0]}: {sr.summary()}",
707
+ analysis_key=f"bootstrap_diff:{names[0]}:{names[1]}",
708
+ correction_family=None,
709
+ raw_reject=sr.reject,
710
+ details={"n": len(common), "strategies": names, **sr.details},
711
+ )
712
+
713
+
714
+ def _tool_mcnemar_evalue(inp: StatsInput, config: dict) -> StatsToolResult:
715
+ """Paired binary comparison of two strategies (McNemar + anytime-valid e-value)."""
716
+ if not inp.groups or len(inp.groups) < 2:
717
+ return StatsToolResult(
718
+ tool="mcnemar_evalue", config=config, ok=False,
719
+ error="need >=2 strategy groups", summary="mcnemar_evalue: <2 groups",
720
+ )
721
+ names = config.get("strategies") or list(inp.groups)[:2]
722
+ names = list(names)[:2]
723
+ common, vecs = _aligned_groups(inp, names)
724
+ if not common:
725
+ return StatsToolResult(
726
+ tool="mcnemar_evalue", config={**config, "strategies": names}, ok=False,
727
+ error="no shared cases between groups", summary="mcnemar_evalue: no overlap",
728
+ )
729
+ a, b = vecs
730
+ sr = compare(
731
+ a, b, paired=True,
732
+ alpha=config.get("alpha", 0.05), min_effect=config.get("min_effect", 0.0),
733
+ )
734
+ return StatsToolResult(
735
+ tool="mcnemar_evalue", config={**config, "strategies": names}, ok=True,
736
+ effect=sr.effect, ci=sr.ci, reject=sr.reject, e_value=sr.e_value,
737
+ p_value=sr.details.get("p_value"), underpowered=sr.underpowered,
738
+ summary=f"{names[1]} vs {names[0]} (paired): {sr.summary()}",
739
+ analysis_key=f"mcnemar_evalue:{names[0]}:{names[1]}",
740
+ correction_family="e_bh",
741
+ raw_reject=sr.reject,
742
+ details={"n": len(common), "strategies": names, **sr.details},
743
+ )
744
+
745
+
746
+ def _tool_friedman_nemenyi(inp: StatsInput, config: dict) -> StatsToolResult:
747
+ """Rank 3+ strategies across shared cases (Friedman omnibus + Nemenyi post-hoc)."""
748
+ if not inp.groups or len(inp.groups) < 3:
749
+ return StatsToolResult(
750
+ tool="friedman_nemenyi", config=config, ok=False,
751
+ error="need >=3 strategy groups", summary="friedman_nemenyi: <3 groups",
752
+ )
753
+ names = list(inp.groups)
754
+ common, vecs = _aligned_groups(inp, names)
755
+ if not common:
756
+ return StatsToolResult(
757
+ tool="friedman_nemenyi", config=config, ok=False,
758
+ error="no shared cases across all groups", summary="friedman_nemenyi: no overlap",
759
+ )
760
+ by_strategy = dict(zip(names, vecs))
761
+ mc = compare_multiple(by_strategy, alpha=config.get("alpha", 0.05))
762
+ return StatsToolResult(
763
+ tool="friedman_nemenyi", config=config, ok=True,
764
+ reject=mc.reject_global, p_value=mc.p_value,
765
+ summary=mc.summary(),
766
+ analysis_key="friedman_nemenyi:global",
767
+ correction_family="bh",
768
+ raw_reject=mc.reject_global,
769
+ details={
770
+ "avg_ranks": mc.avg_ranks,
771
+ "critical_difference": mc.critical_difference,
772
+ "significant_pairs": mc.significant_pairs,
773
+ "n": mc.n,
774
+ },
775
+ )
776
+
777
+
778
+ def _tool_single_rate_evalue(inp: StatsInput, config: dict) -> StatsToolResult:
779
+ """Anytime-valid test that the overall FAIL rate differs from a baseline p0.
780
+
781
+ DESCRIPTIVE ONLY. The result is meaningful only when the case batch is a
782
+ REPRESENTATIVE sample and ``p0`` is a justified baseline (the model's
783
+ natural fail rate on this task). On a curated/enriched batch — the norm for
784
+ diagnosis, where failures are over-sampled so a mechanism can be tested —
785
+ the rate is a sampling artifact and ``p0=0.5`` tests nothing. The verdict
786
+ layer treats this tool as descriptive and never makes it a hypothesis
787
+ headline; pass ``config["p0"]`` = the manifest's recorded base rate to make
788
+ it interpretable.
789
+ """
790
+ if not inp.labels:
791
+ return StatsToolResult(
792
+ tool="single_rate_evalue", config=config, ok=False,
793
+ error="no labeled cases", summary="single_rate_evalue: no labels",
794
+ )
795
+ p0 = config.get("p0", 0.5)
796
+ p0_justified = "p0" in config # explicit baseline vs the meaningless default
797
+ alpha = config.get("alpha", 0.05)
798
+ fails = sum(1 for v in inp.labels.values() if v)
799
+ n = len(inp.labels)
800
+ res = e_value_test(fails, n, p0=p0, alpha=alpha)
801
+ rate = fails / n
802
+ caveat = "" if p0_justified else " (descriptive only: p0=0.5 is not a justified baseline)"
803
+ return StatsToolResult(
804
+ tool="single_rate_evalue", config={**config, "p0": p0}, ok=True,
805
+ # No reported effect when p0 is the unjustified default — its rate − 0.5
806
+ # would otherwise pollute any |effect|-based ranking downstream.
807
+ effect=round(rate - p0, 4) if p0_justified else None,
808
+ e_value=res["e_value"], reject=res["reject"] if p0_justified else False,
809
+ analysis_key="single_rate_evalue:fail_rate",
810
+ correction_family="e_bh" if p0_justified else None,
811
+ raw_reject=res["reject"] if p0_justified else False,
812
+ summary=(
813
+ f"FAIL rate {rate:.1%} ({fails}/{n}) vs p0={p0:.2f}: "
814
+ f"e={res['e_value']:.2f} -> "
815
+ f"{'reject' if (res['reject'] and p0_justified) else 'inconclusive'}{caveat}"
816
+ ),
817
+ details={"fails": fails, "n": n, "rate": round(rate, 4),
818
+ "p0_justified": p0_justified, **res},
819
+ )
820
+
821
+
822
+ def _tool_rank_corr(inp: StatsInput, config: dict) -> StatsToolResult:
823
+ """Kendall τ between a continuous per-case signal and FAIL (monotonic association)."""
824
+ key = config.get("signal") or next(iter(inp.per_case), None)
825
+ cfg = {**config, "signal": key}
826
+ if not key or key not in inp.per_case:
827
+ return StatsToolResult(
828
+ tool="rank_corr", config=cfg, ok=False,
829
+ error="no per-case signal available", summary="rank_corr: no signal",
830
+ )
831
+ sigmap = inp.per_case[key]
832
+ xs: list[float] = []
833
+ ys: list[float] = []
834
+ for cid, is_fail in inp.labels.items():
835
+ if cid in sigmap:
836
+ xs.append(sigmap[cid])
837
+ ys.append(float(is_fail))
838
+ if len(xs) < 3:
839
+ return StatsToolResult(
840
+ tool="rank_corr", config=cfg, ok=False,
841
+ error="need >=3 paired (signal, label) points",
842
+ summary="rank_corr: too few points", details={"n": len(xs)},
843
+ )
844
+ tau = kendall_tau(xs, ys)
845
+ return StatsToolResult(
846
+ tool="rank_corr", config=cfg, ok=True, effect=round(tau, 4),
847
+ summary=f"Kendall τ between '{key}' and FAIL = {tau:+.3f} (n={len(xs)})",
848
+ analysis_key=f"rank_corr:{key}",
849
+ correction_family=None,
850
+ details={"n": len(xs), "tau": round(tau, 4), "signal": key},
851
+ )
852
+
853
+
854
+ # ---------------------------------------------------------------------------
855
+ # Tensor-level omnibus: decode the FAIL label from the full per-case attention
856
+ # map (not a scalar reduction). "Do FAIL and PASS attend differently *anywhere*?"
857
+ # A cross-validated linear decoder's out-of-fold AUC, calibrated by a label-
858
+ # permutation null — valid under the dependence between map cells, and feature-
859
+ # agnostic (robust to which scalar reduction would have mattered). Pure numpy.
860
+ # ---------------------------------------------------------------------------
861
+
862
+ _DECODE_MIN_N = 12 # too few maps to cross-validate a decoder meaningfully
863
+ _DECODE_MIN_PER_CLASS = 3
864
+
865
+
866
+ def _resize2d(m: "np.ndarray", g: int) -> "np.ndarray":
867
+ """Bilinear-resize a 2-D map to ``(g, g)`` (pure numpy; no PIL dep)."""
868
+ h, w = m.shape
869
+ if (h, w) == (g, g):
870
+ return m.astype(np.float64)
871
+ yi = np.linspace(0, h - 1, g)
872
+ xi = np.linspace(0, w - 1, g)
873
+ y0 = np.floor(yi).astype(int)
874
+ x0 = np.floor(xi).astype(int)
875
+ y1 = np.minimum(y0 + 1, h - 1)
876
+ x1 = np.minimum(x0 + 1, w - 1)
877
+ wy = (yi - y0)[:, None]
878
+ wx = (xi - x0)[None, :]
879
+ m = m.astype(np.float64)
880
+ top = m[y0][:, x0] * (1 - wx) + m[y0][:, x1] * wx
881
+ bot = m[y1][:, x0] * (1 - wx) + m[y1][:, x1] * wx
882
+ return top * (1 - wy) + bot * wy
883
+
884
+
885
+ def _cv_oof_scores(X: "np.ndarray", y: "np.ndarray", folds: int, lam: float, seed: int) -> "np.ndarray":
886
+ """Out-of-fold decision scores from a regularized (ridge) linear decoder.
887
+
888
+ Ridge least-squares on ±1 labels — closed-form and stable when features
889
+ outnumber samples (the attention-map regime). Features are standardized on
890
+ each fold's train split; the intercept is dropped (AUC is rank-invariant)."""
891
+ n = len(y)
892
+ rng = np.random.default_rng(seed)
893
+ idx = rng.permutation(n)
894
+ oof = np.zeros(n, dtype=np.float64)
895
+ sizes = np.full(folds, n // folds, dtype=int)
896
+ sizes[: n % folds] += 1
897
+ start = 0
898
+ eye = None
899
+ for fs in sizes:
900
+ te = idx[start:start + fs]
901
+ tr = np.concatenate([idx[:start], idx[start + fs:]])
902
+ start += fs
903
+ if len(tr) < 2 or len(np.unique(y[tr])) < 2:
904
+ continue # degenerate fold → leave OOF scores at 0
905
+ mu = X[tr].mean(0)
906
+ sd = X[tr].std(0) + 1e-8
907
+ xtr = (X[tr] - mu) / sd
908
+ xte = (X[te] - mu) / sd
909
+ if eye is None:
910
+ eye = np.eye(xtr.shape[1])
911
+ yc = 2.0 * y[tr] - 1.0
912
+ w = np.linalg.solve(xtr.T @ xtr + lam * eye, xtr.T @ yc)
913
+ oof[te] = xte @ w
914
+ return oof
915
+
916
+
917
+ def _energy_distance_test(X: "np.ndarray", y: "np.ndarray", *,
918
+ n_perm: int, alpha: float, seed: int) -> "tuple[float, float, bool]":
919
+ """Two-sample ENERGY-DISTANCE permutation test: do the FAIL and PASS rows of
920
+ *X* come from different distributions?
921
+
922
+ E = 2·mean‖x_fail − x_pass‖ − mean‖x_fail − x_fail'‖ − mean‖x_pass − x_pass'‖
923
+ (≥0; larger = more different). More powerful than linear CV-decoding at low n —
924
+ parameter-free and sensitive to nonlinear / higher-moment differences a linear
925
+ boundary misses. The pairwise distance matrix is precomputed ONCE; each
926
+ permutation only re-indexes it (and class sizes are preserved, so the diagonal
927
+ bias cancels), so cost is O(n_perm · n²). Returns ``(energy, perm_p, reject)``."""
928
+ n = len(y)
929
+ sq = (X * X).sum(1)
930
+ d2 = sq[:, None] + sq[None, :] - 2.0 * (X @ X.T)
931
+ np.maximum(d2, 0.0, out=d2)
932
+ dist = np.sqrt(d2)
933
+ yb = np.asarray(y, dtype=bool)
934
+
935
+ def _estat(mask: "np.ndarray") -> float:
936
+ na = int(mask.sum())
937
+ nb = n - na
938
+ if na == 0 or nb == 0:
939
+ return 0.0
940
+ nm = ~mask
941
+ daa = dist[np.ix_(mask, mask)].sum() / (na * na)
942
+ dbb = dist[np.ix_(nm, nm)].sum() / (nb * nb)
943
+ dab = dist[np.ix_(mask, nm)].sum() / (na * nb)
944
+ return 2.0 * dab - daa - dbb
945
+
946
+ obs = _estat(yb)
947
+ rng = np.random.default_rng(seed)
948
+ ge = 1 # +1 (observed) in both numerator and denominator → a valid permutation p
949
+ for _ in range(n_perm):
950
+ if _estat(rng.permutation(yb)) >= obs:
951
+ ge += 1
952
+ p = ge / (n_perm + 1)
953
+ return float(obs), float(p), bool(p < alpha)
954
+
955
+
956
+ def _tool_attention_decoding(inp: StatsInput, config: dict) -> StatsToolResult:
957
+ """Tensor-level omnibus: do FAIL and PASS per-case attention maps differ?
958
+
959
+ Primary test: a two-sample ENERGY-DISTANCE permutation test over the full
960
+ (standardized, resized) maps — parameter-free and more powerful than linear
961
+ decoding at low n. A cross-validated linear-decoder out-of-fold AUC is
962
+ reported alongside as an interpretable (but weaker) companion."""
963
+ key = config.get("signal") or next(iter(inp.per_case_vectors), None)
964
+ cfg = {**config, "signal": key}
965
+ if not key or key not in inp.per_case_vectors:
966
+ return StatsToolResult(
967
+ tool="attention_decoding", config=cfg, ok=False,
968
+ error="no per-case map vectors available", summary="attention_decoding: no maps",
969
+ )
970
+ vecmap = inp.per_case_vectors[key]
971
+ g = int(config.get("grid", 8))
972
+ lam = float(config.get("lam", 1.0))
973
+ n_perm = int(config.get("n_perm", 500)) # tighter p floor (1/(n_perm+1)) than the old 200
974
+ alpha = float(config.get("alpha", 0.05))
975
+ seed = int(config.get("seed", 0))
976
+
977
+ xs: list = []
978
+ ys: list[int] = []
979
+ for cid, is_fail in inp.labels.items():
980
+ m = vecmap.get(cid)
981
+ if m is None:
982
+ continue
983
+ m = np.asarray(m, dtype=np.float64)
984
+ if m.ndim == 1:
985
+ s = int(round(float(np.sqrt(m.size))))
986
+ m = m.reshape(s, s) if s * s == m.size else m.reshape(1, -1)
987
+ if m.ndim != 2 or m.size < 2:
988
+ continue
989
+ xs.append(_resize2d(m, g).ravel())
990
+ ys.append(int(is_fail))
991
+
992
+ n = len(ys)
993
+ n_fail = int(sum(ys))
994
+ if n < _DECODE_MIN_N or n_fail < _DECODE_MIN_PER_CLASS or (n - n_fail) < _DECODE_MIN_PER_CLASS:
995
+ return StatsToolResult(
996
+ tool="attention_decoding", config=cfg, ok=False,
997
+ error=f"insufficient maps for the omnibus (n={n}, fail={n_fail})",
998
+ summary="attention_decoding: underpowered", underpowered=True,
999
+ details={"n": n, "n_fail": n_fail},
1000
+ )
1001
+
1002
+ X = np.vstack(xs)
1003
+ y = np.asarray(ys, dtype=np.float64)
1004
+ # Standardize each map cell so no single high-variance patch dominates the
1005
+ # distance (or the decoder); both tests then see the map's SHAPE, not scale.
1006
+ Xz = (X - X.mean(0)) / (X.std(0) + 1e-8)
1007
+
1008
+ energy, p, reject = _energy_distance_test(Xz, y, n_perm=n_perm, alpha=alpha, seed=seed + 1)
1009
+ folds = max(2, min(int(config.get("folds", 5)), n_fail, n - n_fail))
1010
+ cv_auc = _auc(_cv_oof_scores(Xz, y, folds, lam, seed).tolist(), [int(v) for v in y])
1011
+
1012
+ return StatsToolResult(
1013
+ tool="attention_decoding",
1014
+ config={**cfg, "grid": g, "n_perm": n_perm, "method": "energy_distance"},
1015
+ ok=True, effect=round(float(energy), 4), reject=reject, p_value=round(float(p), 4),
1016
+ underpowered=bool(not reject and p > 0.2 and n < 60),
1017
+ analysis_key=f"attention_decoding:{key}",
1018
+ correction_family="bh",
1019
+ raw_reject=reject,
1020
+ summary=(f"FAIL/PASS attention maps differ: energy-distance={energy:.3f}, "
1021
+ f"permutation p={p:.3f} → {'reject H0 (maps differ)' if reject else 'inconclusive'} "
1022
+ f"(companion CV-AUC={cv_auc:.3f}, n={n})"),
1023
+ details={"n": n, "n_fail": n_fail, "energy_distance": round(float(energy), 4),
1024
+ "perm_p": round(float(p), 4), "cv_auc": round(float(cv_auc), 4),
1025
+ "grid": g, "n_perm": n_perm, "n_features": int(X.shape[1]),
1026
+ "method": "energy_distance"},
1027
+ )
1028
+
1029
+
1030
+ # Registry: name -> callable. Edit STATS_TOOL_CATALOG in lockstep.
1031
+ STATS_TOOLS: dict[str, Callable[[StatsInput, dict], StatsToolResult]] = {
1032
+ "signal_label_assoc": _tool_signal_label_assoc,
1033
+ "bootstrap_diff": _tool_bootstrap_diff,
1034
+ "mcnemar_evalue": _tool_mcnemar_evalue,
1035
+ "friedman_nemenyi": _tool_friedman_nemenyi,
1036
+ "single_rate_evalue": _tool_single_rate_evalue,
1037
+ "rank_corr": _tool_rank_corr,
1038
+ "attention_decoding": _tool_attention_decoding,
1039
+ }
1040
+
1041
+ # Catalog text shown to the LLM selector (name -> when to use it).
1042
+ STATS_TOOL_CATALOG: dict[str, str] = {
1043
+ "signal_label_assoc": (
1044
+ "Unpaired fail-rate difference between cases that exhibit a per-case "
1045
+ "analyzer signal and those that don't (bootstrap CI). Use when you have "
1046
+ "per-case signals AND PASS/FAIL labels — this is the main M2 test."
1047
+ ),
1048
+ "bootstrap_diff": (
1049
+ "Unpaired fail-rate difference between two strategy groups (bootstrap CI). "
1050
+ "Needs >=2 strategy groups (findings['by_strategy'])."
1051
+ ),
1052
+ "mcnemar_evalue": (
1053
+ "Paired binary comparison of two strategies on the same cases "
1054
+ "(McNemar + anytime-valid e-value). Needs exactly 2 paired strategy groups."
1055
+ ),
1056
+ "friedman_nemenyi": (
1057
+ "Rank 3+ strategies across shared cases (Friedman omnibus + Nemenyi "
1058
+ "post-hoc). Needs >=3 strategy groups."
1059
+ ),
1060
+ "single_rate_evalue": (
1061
+ "DESCRIPTIVE context only: tests whether the overall FAIL rate differs "
1062
+ "from a baseline p0. Meaningful ONLY on a representative sample with a "
1063
+ "justified p0 (pass config['p0'] = the natural base rate); on a curated/"
1064
+ "enriched diagnosis batch it tests nothing. Never decides a hypothesis."
1065
+ ),
1066
+ "rank_corr": (
1067
+ "Kendall tau between a continuous per-case signal and FAIL (monotonic "
1068
+ "association). Needs a continuous per-case signal."
1069
+ ),
1070
+ "attention_decoding": (
1071
+ "Tensor-level OMNIBUS: a two-sample ENERGY-DISTANCE permutation test over "
1072
+ "the FULL per-case attention map (not a scalar reduction), with a CV "
1073
+ "linear-decoder AUC reported alongside. Answers 'do FAIL and PASS attend "
1074
+ "differently anywhere?' — feature-agnostic and sensitive to nonlinear / "
1075
+ "distributional differences. Needs per-case map vectors (findings carry "
1076
+ "the scalars; the maps come from artifacts['per_case_maps'])."
1077
+ ),
1078
+ }
1079
+
1080
+
1081
+ def run_stats_tool(name: str, inp: StatsInput, config: dict | None = None) -> StatsToolResult:
1082
+ """Run a single catalog tool by name. Raises KeyError for unknown names."""
1083
+ tool = STATS_TOOLS[name]
1084
+ return tool(inp, config or {})
1085
+
1086
+
1087
+ # ---------------------------------------------------------------------------
1088
+ # Deterministic planner (fallback when no judge / LLM selection fails)
1089
+ # ---------------------------------------------------------------------------
1090
+
1091
+ def default_plan(
1092
+ inp: StatsInput,
1093
+ max_signals: int | None = None,
1094
+ ) -> list[tuple[str, dict, str]]:
1095
+ """Deterministic ``[(tool, config, rationale)]`` plan from the data shape.
1096
+
1097
+ The implementation delegates to the generic M2 planner so per-case signals
1098
+ are ranked by testability instead of original column order. ``max_signals``
1099
+ remains for backward compatibility; ``None`` means test every ranked signal.
1100
+ """
1101
+ from evalrx.analysis.planner import plan_stats_input
1102
+
1103
+ return [item.as_legacy_tuple() for item in plan_stats_input(inp, max_signals=max_signals)]
1104
+
1105
+
1106
+ # ---------------------------------------------------------------------------
1107
+ # Multiple-testing correction + visualization
1108
+ # ---------------------------------------------------------------------------
1109
+
1110
+ def fdr_correct(results: list[StatsToolResult], alpha: float = 0.05) -> dict[str, Any]:
1111
+ """Apply multiplicity correction across all supported result families.
1112
+
1113
+ e-values use e-BH; p-values use BH. The returned ``rejected_tools`` field is
1114
+ preserved for existing M1-M4 consumers, while ``rejected_result_keys`` and
1115
+ ``families`` expose the precise generalized-M2 family membership.
1116
+ """
1117
+ return correct_results(results, alpha=alpha)
1118
+
1119
+
1120
+ def plot_effects(results: list[StatsToolResult], out_path: str) -> str | None:
1121
+ """Forest plot of effect ± CI for tools that produced both. Returns path or None."""
1122
+ items = [
1123
+ (r.tool, r.effect, r.ci)
1124
+ for r in results
1125
+ if r.ok and r.effect is not None and r.ci is not None
1126
+ ]
1127
+ if not items:
1128
+ return None
1129
+ try:
1130
+ import matplotlib
1131
+ matplotlib.use("Agg")
1132
+ import matplotlib.pyplot as plt
1133
+ except Exception: # pragma: no cover - matplotlib optional
1134
+ return None
1135
+
1136
+ os.makedirs(os.path.dirname(out_path) or ".", exist_ok=True)
1137
+ labels = [t for t, _, _ in items]
1138
+ effects = [e for _, e, _ in items]
1139
+ lows = [e - ci[0] for _, e, ci in items]
1140
+ highs = [ci[1] - e for _, e, ci in items]
1141
+ ys = list(range(len(items)))
1142
+
1143
+ fig, ax = plt.subplots(figsize=(7, 0.6 * len(items) + 1.5))
1144
+ ax.errorbar(effects, ys, xerr=[lows, highs], fmt="o", capsize=4, color="#2b6cb0")
1145
+ ax.axvline(0.0, color="grey", linestyle="--", linewidth=1)
1146
+ ax.set_yticks(ys)
1147
+ ax.set_yticklabels(labels)
1148
+ ax.set_xlabel("effect size (fail-rate difference)")
1149
+ ax.set_title("M2 statistical effects (± CI)")
1150
+ fig.tight_layout()
1151
+ fig.savefig(out_path, dpi=120)
1152
+ plt.close(fig)
1153
+ return out_path