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.
- evalrx/__init__.py +139 -0
- evalrx/agent_assets/__init__.py +2 -0
- evalrx/agent_assets/skills/README.md +28 -0
- evalrx/agent_assets/skills/eval-chart-style/SKILL.md +172 -0
- evalrx/agent_assets/skills/evalrx-report-ui/SKILL.md +116 -0
- evalrx/agent_assets/skills/nature-figure/LICENSE +201 -0
- evalrx/agent_assets/skills/nature-figure/README.md +412 -0
- evalrx/agent_assets/skills/nature-figure/SKILL.md +60 -0
- evalrx/agent_assets/skills/nature-figure/manifest.yaml +59 -0
- evalrx/agent_assets/skills/nature-figure/references/api.md +436 -0
- evalrx/agent_assets/skills/nature-figure/references/backend-selection.md +100 -0
- evalrx/agent_assets/skills/nature-figure/references/chart-types.md +281 -0
- evalrx/agent_assets/skills/nature-figure/references/common-patterns.md +350 -0
- evalrx/agent_assets/skills/nature-figure/references/demos.md +65 -0
- evalrx/agent_assets/skills/nature-figure/references/design-theory.md +439 -0
- evalrx/agent_assets/skills/nature-figure/references/figure-contract.md +93 -0
- evalrx/agent_assets/skills/nature-figure/references/figure-legend-conventions.md +71 -0
- evalrx/agent_assets/skills/nature-figure/references/nature-2026-observations.md +112 -0
- evalrx/agent_assets/skills/nature-figure/references/qa-contract.md +119 -0
- evalrx/agent_assets/skills/nature-figure/references/r-template-index.md +66 -0
- evalrx/agent_assets/skills/nature-figure/references/r-workflow.md +161 -0
- evalrx/agent_assets/skills/nature-figure/references/tutorials.md +251 -0
- evalrx/agent_assets/skills/nature-figure/static/core/contract.md +29 -0
- evalrx/agent_assets/skills/nature-figure/static/core/stance.md +37 -0
- evalrx/agent_assets/skills/nature-figure/static/fragments/backend/python.md +37 -0
- evalrx/agent_assets/skills/nature-figure/static/fragments/backend/r.md +44 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/SKILL.md +213 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/assets/analysis_report_template.md +53 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/references/model_selection.md +72 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/explanatory_var_eda.R +130 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/explanatory_var_eda.py +150 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/fit_outcome_model.R +181 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/fit_outcome_model.py +186 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/univariate_eda.R +149 -0
- evalrx/agent_assets/skills/outcome-driver-analysis/scripts/univariate_eda.py +177 -0
- evalrx/agent_assets/skills.py +27 -0
- evalrx/agent_runtime/__init__.py +78 -0
- evalrx/agent_runtime/_docker_runner.py +89 -0
- evalrx/agent_runtime/cli_runtime.py +103 -0
- evalrx/agent_runtime/cli_transcript.py +138 -0
- evalrx/agent_runtime/cli_types.py +68 -0
- evalrx/agent_runtime/codegen/__init__.py +5 -0
- evalrx/agent_runtime/codegen/runner.py +94 -0
- evalrx/agent_runtime/experiment_harness.py +117 -0
- evalrx/agent_runtime/factory.py +102 -0
- evalrx/agent_runtime/json_shape.py +44 -0
- evalrx/agent_runtime/judges/__init__.py +28 -0
- evalrx/agent_runtime/judges/agy.py +179 -0
- evalrx/agent_runtime/judges/autodetect.py +135 -0
- evalrx/agent_runtime/judges/claude.py +159 -0
- evalrx/agent_runtime/judges/codex.py +120 -0
- evalrx/agent_runtime/providers/__init__.py +21 -0
- evalrx/agent_runtime/providers/antigravity.py +31 -0
- evalrx/agent_runtime/providers/base.py +145 -0
- evalrx/agent_runtime/providers/claude_code.py +49 -0
- evalrx/agent_runtime/providers/codex.py +37 -0
- evalrx/agent_runtime/providers/gemini_cli.py +26 -0
- evalrx/agent_runtime/providers/kimi_cli.py +27 -0
- evalrx/agent_runtime/providers/opencode.py +27 -0
- evalrx/agent_runtime/providers/registry.py +58 -0
- evalrx/agent_runtime/sandbox.py +517 -0
- evalrx/agent_runtime/skill_audit.py +143 -0
- evalrx/agent_runtime/skills/__init__.py +19 -0
- evalrx/agent_runtime/skills/installer.py +68 -0
- evalrx/agent_runtime/skills/prompt_policy.py +86 -0
- evalrx/agent_runtime/skills/resolver.py +19 -0
- evalrx/analysis/__init__.py +132 -0
- evalrx/analysis/adjudicate.py +154 -0
- evalrx/analysis/analysis_module.py +361 -0
- evalrx/analysis/api.py +171 -0
- evalrx/analysis/case_studio.py +651 -0
- evalrx/analysis/cli.py +114 -0
- evalrx/analysis/dashboard.py +350 -0
- evalrx/analysis/eval_case_matrix.py +118 -0
- evalrx/analysis/eval_viz_theme.py +833 -0
- evalrx/analysis/explore_run.py +333 -0
- evalrx/analysis/explorer.py +1276 -0
- evalrx/analysis/failure_modes.py +607 -0
- evalrx/analysis/fused_pipeline.py +489 -0
- evalrx/analysis/holdout.py +300 -0
- evalrx/analysis/hypothesis_agent.py +230 -0
- evalrx/analysis/narration.py +177 -0
- evalrx/analysis/operationalize.py +442 -0
- evalrx/analysis/plain_language.py +42 -0
- evalrx/analysis/planner.py +283 -0
- evalrx/analysis/probe_search.py +203 -0
- evalrx/analysis/profile.py +268 -0
- evalrx/analysis/prompts/__init__.py +0 -0
- evalrx/analysis/prompts/explorer.py +417 -0
- evalrx/analysis/prompts/failure_modes.py +33 -0
- evalrx/analysis/prompts/holdout.py +27 -0
- evalrx/analysis/prompts/hypothesis_agent.py +78 -0
- evalrx/analysis/prompts/run_codebase.py +47 -0
- evalrx/analysis/prompts/stats_agent.py +72 -0
- evalrx/analysis/prompts/stats_tool_generator.py +43 -0
- evalrx/analysis/result_marker.py +47 -0
- evalrx/analysis/run_codebase.py +242 -0
- evalrx/analysis/run_view.py +205 -0
- evalrx/analysis/stage_views.py +93 -0
- evalrx/analysis/stats_agent.py +944 -0
- evalrx/analysis/stats_tool_agent.py +261 -0
- evalrx/analysis/stats_tool_generator.py +415 -0
- evalrx/analysis/stats_tools.py +1153 -0
- evalrx/analysis/trajectory_records.py +193 -0
- evalrx/analysis/workbench.py +431 -0
- evalrx/analyzers/__init__.py +42 -0
- evalrx/analyzers/agent/__init__.py +25 -0
- evalrx/analyzers/agent/counterfactual.py +84 -0
- evalrx/analyzers/agent/first_error_judge.py +96 -0
- evalrx/analyzers/agent/ignored_obs.py +81 -0
- evalrx/analyzers/agent/loop_detect.py +79 -0
- evalrx/analyzers/agent/reliability.py +165 -0
- evalrx/analyzers/agent/tool_shap.py +225 -0
- evalrx/analyzers/agent/trajectory_rubric.py +168 -0
- evalrx/analyzers/attention/__init__.py +19 -0
- evalrx/analyzers/attention/relative_attn.py +610 -0
- evalrx/analyzers/attention/rollout.py +73 -0
- evalrx/analyzers/attention/sink.py +56 -0
- evalrx/analyzers/attention/summary.py +190 -0
- evalrx/analyzers/attribution/__init__.py +6 -0
- evalrx/analyzers/attribution/generic_attn.py +31 -0
- evalrx/analyzers/attribution/gradcam.py +30 -0
- evalrx/analyzers/base.py +12 -0
- evalrx/analyzers/geometry/__init__.py +6 -0
- evalrx/analyzers/geometry/cka.py +70 -0
- evalrx/analyzers/geometry/linear_probe.py +157 -0
- evalrx/analyzers/hallucination/__init__.py +9 -0
- evalrx/analyzers/hallucination/chair.py +78 -0
- evalrx/analyzers/hallucination/opera.py +29 -0
- evalrx/analyzers/hallucination/pope.py +119 -0
- evalrx/analyzers/hallucination/selfcheck.py +155 -0
- evalrx/analyzers/hallucination/vcd.py +29 -0
- evalrx/analyzers/lens/__init__.py +7 -0
- evalrx/analyzers/lens/layer_contrast.py +133 -0
- evalrx/analyzers/lens/logit_lens.py +138 -0
- evalrx/analyzers/lens/tuned_lens.py +30 -0
- evalrx/analyzers/patching/__init__.py +5 -0
- evalrx/analyzers/patching/causal_trace.py +30 -0
- evalrx/analyzers/perturbation/__init__.py +23 -0
- evalrx/analyzers/perturbation/_shapley.py +54 -0
- evalrx/analyzers/perturbation/context_shap.py +174 -0
- evalrx/analyzers/perturbation/cot_faithfulness.py +239 -0
- evalrx/analyzers/perturbation/format_sensitivity.py +237 -0
- evalrx/analyzers/perturbation/mm_shap.py +146 -0
- evalrx/analyzers/perturbation/modality_ablation.py +196 -0
- evalrx/analyzers/perturbation/perturbation_battery.py +274 -0
- evalrx/analyzers/perturbation/prompt_contrast.py +265 -0
- evalrx/analyzers/perturbation/rise.py +94 -0
- evalrx/analyzers/perturbation/vl_shap.py +102 -0
- evalrx/analyzers/reasoning/__init__.py +33 -0
- evalrx/analyzers/reasoning/_text.py +328 -0
- evalrx/analyzers/reasoning/answer_extraction_audit.py +327 -0
- evalrx/analyzers/reasoning/arith_audit.py +226 -0
- evalrx/analyzers/reasoning/contamination.py +214 -0
- evalrx/analyzers/reasoning/knowledge_split.py +253 -0
- evalrx/analyzers/reasoning/self_repair.py +246 -0
- evalrx/analyzers/reasoning/step_rollout_value.py +216 -0
- evalrx/analyzers/reasoning/termination_audit.py +258 -0
- evalrx/analyzers/uncertainty/__init__.py +18 -0
- evalrx/analyzers/uncertainty/calibration.py +174 -0
- evalrx/analyzers/uncertainty/coverage_gap.py +199 -0
- evalrx/analyzers/uncertainty/entropy.py +90 -0
- evalrx/analyzers/uncertainty/logprob_entropy.py +69 -0
- evalrx/analyzers/uncertainty/self_consistency.py +204 -0
- evalrx/analyzers/uncertainty/verbalized_conf.py +64 -0
- evalrx/cli.py +411 -0
- evalrx/config.py +77 -0
- evalrx/contract/__init__.py +179 -0
- evalrx/contract/common.py +452 -0
- evalrx/contract/emit.py +948 -0
- evalrx/contract/export.py +237 -0
- evalrx/contract/m1.py +325 -0
- evalrx/contract/m2.py +317 -0
- evalrx/contract/m3.py +165 -0
- evalrx/contract/m4.py +130 -0
- evalrx/contract/m5.py +292 -0
- evalrx/contract/methodology.py +76 -0
- evalrx/contract/pre_m1.py +58 -0
- evalrx/contract/typescript.py +140 -0
- evalrx/core/__init__.py +85 -0
- evalrx/core/analyzer.py +174 -0
- evalrx/core/capability.py +54 -0
- evalrx/core/case.py +443 -0
- evalrx/core/experiment.py +106 -0
- evalrx/core/model.py +198 -0
- evalrx/core/pipeline.py +42 -0
- evalrx/core/registry.py +142 -0
- evalrx/core/result.py +64 -0
- evalrx/core/spec.py +173 -0
- evalrx/core/tokentype.py +165 -0
- evalrx/core/tool.py +92 -0
- evalrx/datasets/__init__.py +41 -0
- evalrx/datasets/base.py +68 -0
- evalrx/datasets/gui_os.py +52 -0
- evalrx/datasets/llm_qa.py +57 -0
- evalrx/datasets/pure_qa.py +12 -0
- evalrx/datasets/vlm_qa.py +695 -0
- evalrx/datasets/web_search_qa.py +52 -0
- evalrx/eval_agent/__init__.py +341 -0
- evalrx/eval_agent/_tools.py +81 -0
- evalrx/eval_agent/ab_runner.py +50 -0
- evalrx/eval_agent/agentic/__init__.py +43 -0
- evalrx/eval_agent/agentic/actions.py +216 -0
- evalrx/eval_agent/agentic/board.py +107 -0
- evalrx/eval_agent/agentic/loop.py +190 -0
- evalrx/eval_agent/agentic/tools.py +538 -0
- evalrx/eval_agent/checkpoint.py +57 -0
- evalrx/eval_agent/cli_agent.py +59 -0
- evalrx/eval_agent/cli_skills.py +5 -0
- evalrx/eval_agent/evolution.py +396 -0
- evalrx/eval_agent/git_manager.py +215 -0
- evalrx/eval_agent/hypothesis.py +172 -0
- evalrx/eval_agent/label_quarantine.py +209 -0
- evalrx/eval_agent/legacy.py +530 -0
- evalrx/eval_agent/log_schema.py +497 -0
- evalrx/eval_agent/loop.py +2159 -0
- evalrx/eval_agent/loop_reports.py +116 -0
- evalrx/eval_agent/model_instrumentation.py +282 -0
- evalrx/eval_agent/narration.py +193 -0
- evalrx/eval_agent/nl_runner.py +460 -0
- evalrx/eval_agent/orchestrator.py +61 -0
- evalrx/eval_agent/preregister.py +93 -0
- evalrx/eval_agent/prompts/__init__.py +1 -0
- evalrx/eval_agent/prompts/agentic.py +46 -0
- evalrx/eval_agent/prompts/case_discovery.py +25 -0
- evalrx/eval_agent/prompts/diagnosis.py +125 -0
- evalrx/eval_agent/prompts/experiment_writer.py +265 -0
- evalrx/eval_agent/prompts/explore_step.py +37 -0
- evalrx/eval_agent/prompts/fix_agent.py +257 -0
- evalrx/eval_agent/prompts/hypothesis_tester.py +15 -0
- evalrx/eval_agent/prompts/nl_runner.py +38 -0
- evalrx/eval_agent/prompts/probe_agent.py +25 -0
- evalrx/eval_agent/prompts/probe_candidate_generator.py +14 -0
- evalrx/eval_agent/prompts/probe_generator.py +35 -0
- evalrx/eval_agent/prompts/whitebox_probe_generator.py +38 -0
- evalrx/eval_agent/report.py +58 -0
- evalrx/eval_agent/run_context.py +354 -0
- evalrx/eval_agent/run_log.schema.json +1215 -0
- evalrx/eval_agent/run_logger_v2.py +1764 -0
- evalrx/eval_agent/run_metadata.py +208 -0
- evalrx/eval_agent/stages/__init__.py +56 -0
- evalrx/eval_agent/stages/case_discovery.py +293 -0
- evalrx/eval_agent/stages/diagnosis.py +1017 -0
- evalrx/eval_agent/stages/experiment_writer.py +1634 -0
- evalrx/eval_agent/stages/fix_agent.py +3916 -0
- evalrx/eval_agent/stages/fix_internals.py +499 -0
- evalrx/eval_agent/stages/fix_pipeline.py +725 -0
- evalrx/eval_agent/stages/fix_tiers.py +187 -0
- evalrx/eval_agent/stages/fix_tools.py +1034 -0
- evalrx/eval_agent/stages/hypothesis_tester.py +1014 -0
- evalrx/eval_agent/stages/probe.py +439 -0
- evalrx/eval_agent/stages/probe_agent.py +1079 -0
- evalrx/eval_agent/stages/probe_candidate_generator.py +128 -0
- evalrx/eval_agent/stages/probe_generator.py +326 -0
- evalrx/eval_agent/stages/probe_search_agent.py +106 -0
- evalrx/eval_agent/stages/protocol.py +112 -0
- evalrx/eval_agent/stages/repair_catalog.py +273 -0
- evalrx/eval_agent/stages/surgery.py +524 -0
- evalrx/eval_agent/stages/whitebox_probe_generator.py +351 -0
- evalrx/eval_agent/store.py +231 -0
- evalrx/logging_utils.py +112 -0
- evalrx/models/__init__.py +161 -0
- evalrx/models/_discover.py +101 -0
- evalrx/models/agent.py +380 -0
- evalrx/models/backends/__init__.py +58 -0
- evalrx/models/backends/api.py +169 -0
- evalrx/models/backends/base.py +57 -0
- evalrx/models/backends/gemini_compat.py +579 -0
- evalrx/models/backends/hf_local.py +2074 -0
- evalrx/models/backends/openai_compat.py +301 -0
- evalrx/models/backends/vllm_offline.py +116 -0
- evalrx/models/base.py +24 -0
- evalrx/models/blackbox/__init__.py +4 -0
- evalrx/models/blackbox/agent.py +31 -0
- evalrx/models/blackbox/base.py +29 -0
- evalrx/models/blackbox/gemini.py +279 -0
- evalrx/models/blackbox/llm_api.py +17 -0
- evalrx/models/blackbox/vlm_api.py +17 -0
- evalrx/models/compose.py +66 -0
- evalrx/models/inference.py +88 -0
- evalrx/models/paper_methods/__init__.py +8 -0
- evalrx/models/paper_methods/aad.py +53 -0
- evalrx/models/paper_methods/ifcd.py +204 -0
- evalrx/models/paper_methods/pai.py +164 -0
- evalrx/models/paper_methods/tcd.py +202 -0
- evalrx/models/paper_methods/vcd.py +45 -0
- evalrx/models/paper_methods/vicrop.py +137 -0
- evalrx/models/toolcodec.py +143 -0
- evalrx/models/tools/__init__.py +20 -0
- evalrx/models/tools/perception.py +300 -0
- evalrx/models/tools/visual.py +174 -0
- evalrx/models/whitebox/__init__.py +26 -0
- evalrx/models/whitebox/agent.py +31 -0
- evalrx/models/whitebox/base.py +24 -0
- evalrx/models/whitebox/qwen.py +61 -0
- evalrx/models/whitebox/qwen2_5_omni.py +29 -0
- evalrx/models/whitebox/qwen2_audio.py +25 -0
- evalrx/models/whitebox/qwen_omni.py +53 -0
- evalrx/models/whitebox/qwen_vl.py +62 -0
- evalrx/observability/__init__.py +21 -0
- evalrx/observability/envelope.py +122 -0
- evalrx/observability/outbox.py +111 -0
- evalrx/observability/tracer.py +882 -0
- evalrx/reporting/__init__.py +28 -0
- evalrx/reporting/case_study.py +947 -0
- evalrx/reporting/compiler.py +587 -0
- evalrx/reporting/dynamic.py +1882 -0
- evalrx/reporting/html_report.py +2225 -0
- evalrx/reporting/langfuse_exporter.py +38 -0
- evalrx/reporting/langfuse_source.py +155 -0
- evalrx/reporting/model.py +151 -0
- evalrx/reporting/run_events.py +184 -0
- evalrx/reporting/server.py +557 -0
- evalrx/reporting/stages.py +58 -0
- evalrx/reporting/static_export.py +142 -0
- evalrx/reporting/web_dist/index.html +146 -0
- evalrx/specs.py +727 -0
- evalrx/stats/__init__.py +47 -0
- evalrx/stats/api.py +192 -0
- evalrx/stats/bootstrap.py +86 -0
- evalrx/stats/ebh.py +27 -0
- evalrx/stats/evalue.py +98 -0
- evalrx/stats/friedman.py +138 -0
- evalrx/stats/mcnemar.py +40 -0
- evalrx/stats/multiplicity.py +159 -0
- evalrx/stats/subset_sampling.py +55 -0
- evalrx/term_links.py +43 -0
- evalrx/viz/__init__.py +7 -0
- evalrx/viz/labels.py +77 -0
- evalrx/viz/prompts.py +39 -0
- evalrx/viz/renderer.py +590 -0
- evalrx/viz/schema.py +36 -0
- evalrx/viz/style.py +134 -0
- evalrx-0.1.2.dist-info/METADATA +532 -0
- evalrx-0.1.2.dist-info/RECORD +339 -0
- evalrx-0.1.2.dist-info/WHEEL +5 -0
- evalrx-0.1.2.dist-info/entry_points.txt +3 -0
- evalrx-0.1.2.dist-info/licenses/LICENSE +121 -0
- 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
|