SabaPivot's picture
Upgrade all claim evidence using high-scoring peer protocols with attribution
854f51d verified
Raw History Blame Contribute Delete
17.4 kB
#!/usr/bin/env python3
"""Direct Theorem 3.1 limit audit on every released DCF MATH graph.
This runner evaluates the theorem's literal coupled schedule, not the practical
finite-temperature implementation. Risk scores and graph data come from the
pinned official release; the limit expressions below are transcribed from the
paper's Theorem 3.1 and evaluated independently in float64 log space.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import math
import sys
import types
from pathlib import Path
import numpy as np
import torch
PAPER_ID = "XfndtVLIub"
OFFICIAL_COMMIT = "0b4d5487a9868a18c4f9aa5b3d96cdccc705ca97"
DATA_SHA256 = "2711d37c337c552396772794aface3e76490a98af3132e4ae14654a779ec6596"
TEMPERATURES = (0.5, 0.2, 0.1, 0.05, 0.02, 0.01, 0.005, 0.002, 0.001)
ALPHAS = tuple(i / 100.0 for i in range(1, 16))
FINAL_TOLERANCE = 1e-10
def sha256(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest()
def write_json(path: Path, value: object) -> None:
path.write_text(json.dumps(value, indent=2, sort_keys=True, allow_nan=False) + "\n", encoding="utf-8")
def write_csv(path: Path, fieldnames: list[str], rows: list[dict[str, object]]) -> None:
with path.open("w", encoding="utf-8", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=fieldnames, lineterminator="\n")
writer.writeheader()
writer.writerows(rows)
def log_sigmoid(value: np.ndarray) -> np.ndarray:
"""Stable log(sigmoid(value))."""
return -np.logaddexp(0.0, -value)
def log1mexp(log_x: np.ndarray) -> np.ndarray:
"""Stable log(1-exp(log_x)) for log_x <= 0."""
clipped = np.minimum(log_x, 0.0)
cutoff = -math.log(2.0)
result = np.empty_like(clipped)
low = clipped < cutoff
exact_zero = clipped == 0.0
middle = ~(low | exact_zero)
result[low] = np.log1p(-np.exp(clipped[low]))
result[middle] = np.log(-np.expm1(clipped[middle]))
result[exact_zero] = -np.inf
return result
def softmax(value: np.ndarray) -> np.ndarray:
shifted = value - np.max(value)
weights = np.exp(shifted)
return weights / np.sum(weights)
def threshold_grid(risk: np.ndarray) -> np.ndarray:
return np.concatenate(([float(np.min(risk) - 1.0)], np.unique(risk), [float(np.max(risk) + 1.0)]))
def hard_score(risk: np.ndarray, labels: np.ndarray, ancestors: np.ndarray, grid: np.ndarray) -> float:
false_nodes = np.flatnonzero(labels == 0)
safe: list[float] = []
for tau in grid:
coherent_false_selected = False
for node in false_nodes:
required = ancestors[:, node].copy()
required[node] = True
if bool(np.all(risk[required] <= tau)):
coherent_false_selected = True
break
if not coherent_false_selected:
safe.append(float(tau))
if not safe:
raise RuntimeError("hard CF grid has no safe threshold")
return max(safe)
def soft_score(
risk: np.ndarray,
labels: np.ndarray,
ancestors: np.ndarray,
grid: np.ndarray,
temperature: float,
*,
beta_mode: str = "coupled",
violation_mode: str = "theorem",
) -> tuple[float, float, float, float]:
"""Evaluate Theorem 3.1's score and return score, beta, tau_s, lambda."""
span = float(grid[-1] - grid[0])
if not span > 0.0:
raise RuntimeError("degenerate threshold grid")
lambda_ = 0.5 / span
tau_s = temperature ** 0.5
beta = temperature ** -1.0 if beta_mode == "coupled" else 8.0
# The theorem requires the +sqrt(T_p) margin. Columns correspond to grid
# thresholds and rows to released claim nodes.
logits = (grid[None, :] - risk[:, None] + math.sqrt(temperature)) / temperature
log_keep = log_sigmoid(logits)
false_nodes = np.flatnonzero(labels == 0)
if false_nodes.size == 0:
log_q_global = np.zeros(grid.size, dtype=np.float64)
else:
negative_terms: list[np.ndarray] = []
for node in false_nodes:
required = ancestors[:, node].copy()
required[node] = True
log_coherent = np.mean(log_keep[required, :], axis=0)
negative_terms.append(log1mexp(log_coherent))
log_q_global = np.mean(np.stack(negative_terms, axis=0), axis=0)
# Q_tau^(1/tau_s) is the theorem's sharpened validity object.
sharpened_validity = np.exp(log_q_global / tau_s)
violation = 1.0 - sharpened_validity
if violation_mode == "none":
violation = np.zeros_like(violation)
elif violation_mode != "theorem":
raise ValueError(violation_mode)
objective = lambda_ * grid - violation
weights = softmax(beta * objective)
score = float(np.dot(weights, grid))
if not (math.isfinite(score) and np.all(np.isfinite(weights))):
raise RuntimeError("non-finite theorem evaluation")
return score, beta, tau_s, lambda_
def conformal_quantile(scores: list[float], alpha: float) -> float:
ordered = sorted(scores)
# Standard split-conformal upper order statistic, clipped at n.
rank = min(len(ordered), int(math.ceil((len(ordered) + 1) * (1.0 - alpha))))
return float(ordered[rank - 1])
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--source-root", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, required=True)
args = parser.parse_args()
root = args.source_root.resolve()
source = root / "source_current"
theorem_path = root / "source_extract/hard_recovery.txt"
data_path = source / "data/MATH_open_subclaims_with_scores_and_semantic_eval.json"
args.output_dir.mkdir(parents=True, exist_ok=True)
theorem_text = theorem_path.read_text(encoding="utf-8")
theorem_markers = {
"sqrt_margin": r"+ \sqrt{T_p}",
"lambda_grid_span": r"\sup(\lambda\mathcal{T})-\inf(\lambda\mathcal{T}) \leq 1",
"single_limit_schedule": r"Setting $\tau_s = T_p^{s}$ and $\beta = T_p^a$",
"score_conclusion": r"soft nonconformity score from CF",
}
marker_checks = {name: marker in theorem_text for name, marker in theorem_markers.items()}
# The conclusion is worded as "hard nonconformity score from CF" in source.
marker_checks["score_conclusion"] = "hard nonconformity score from CF" in theorem_text
if not all(marker_checks.values()):
raise RuntimeError({"theorem_source_marker_drift": marker_checks})
if sha256(data_path) != DATA_SHA256:
raise RuntimeError("released dataset SHA-256 drift")
torch.manual_seed(260420098)
torch.use_deterministic_algorithms(True)
torchsort = types.ModuleType("torchsort")
torchsort.soft_sort = lambda values, regularization_strength=1e-4: torch.sort(values, dim=-1).values
sys.modules.setdefault("torchsort", torchsort)
sys.path.insert(0, str(source))
from src.differentiable_conformal_factuality import compute_risk # type: ignore
from src.models import ForwardScorer # type: ignore
from src.reasonining_graph_dataset import Reasoning_Graph_Dataset # type: ignore
dataset = Reasoning_Graph_Dataset(str(data_path), ["frequency-score"])
scorer = ForwardScorer(0)
records: list[tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, float]] = []
total_nodes = 0
total_edges = 0
false_graphs = 0
for graph, labels_t in dataset[:]:
risks_t = compute_risk(scorer(graph["features"]), graph["adj"], C=6.0, beta_mix=0.5, scalar_noise=0.0)
risk = risks_t.detach().to(dtype=torch.float64).cpu().numpy()
labels = labels_t.detach().cpu().numpy()
ancestors = graph["ancestors"].detach().cpu().numpy().astype(bool)
grid = threshold_grid(risk)
hard = hard_score(risk, labels, ancestors, grid)
records.append((risk, labels, ancestors, grid, hard))
total_nodes += int(labels.size)
total_edges += int(graph["adj"].sum().item())
false_graphs += int(np.any(labels == 0))
if len(records) != 50 or total_nodes != 503:
raise RuntimeError("released-data scale drift")
path_rows: list[dict[str, object]] = []
errors_by_temperature: dict[float, list[float]] = {}
soft_scores_by_temperature: dict[float, list[float]] = {}
lambda_spans: list[float] = []
for temperature in TEMPERATURES:
errors: list[float] = []
soft_scores: list[float] = []
for graph_index, (risk, labels, ancestors, grid, hard) in enumerate(records):
soft, beta, tau_s, lambda_ = soft_score(risk, labels, ancestors, grid, temperature)
error = abs(soft - hard)
errors.append(error)
soft_scores.append(soft)
lambda_spans.append(lambda_ * float(grid[-1] - grid[0]))
path_rows.append({
"graph_index": graph_index,
"nodes": int(labels.size),
"false_nodes": int(np.sum(labels == 0)),
"temperature": format(temperature, ".17g"),
"tau_s": format(tau_s, ".17g"),
"beta": format(beta, ".17g"),
"lambda": format(lambda_, ".17g"),
"lambda_grid_span": format(lambda_ * float(grid[-1] - grid[0]), ".17g"),
"hard_score": format(hard, ".17g"),
"soft_score": format(soft, ".17g"),
"absolute_error": format(error, ".17g"),
"within_final_tolerance": str(error <= FINAL_TOLERANCE).lower(),
})
errors_by_temperature[temperature] = errors
soft_scores_by_temperature[temperature] = soft_scores
hard_scores = [record[-1] for record in records]
quantile_rows: list[dict[str, object]] = []
quantile_bound_holds = True
for temperature in TEMPERATURES:
uniform_bound = max(errors_by_temperature[temperature])
for alpha in ALPHAS:
hard_q = conformal_quantile(hard_scores, alpha)
soft_q = conformal_quantile(soft_scores_by_temperature[temperature], alpha)
q_error = abs(soft_q - hard_q)
bound_holds = q_error <= uniform_bound + 1e-14
quantile_bound_holds = quantile_bound_holds and bound_holds
quantile_rows.append({
"temperature": format(temperature, ".17g"),
"alpha": format(alpha, ".17g"),
"hard_quantile": format(hard_q, ".17g"),
"soft_quantile": format(soft_q, ".17g"),
"absolute_error": format(q_error, ".17g"),
"uniform_score_error_bound": format(uniform_bound, ".17g"),
"order_statistic_bound_holds": str(bound_holds).lower(),
})
control_rows: list[dict[str, object]] = []
control_summaries: dict[str, dict[str, object]] = {}
final_temperature = TEMPERATURES[-1]
for control_name, beta_mode, violation_mode in (
("fixed_beta_8", "fixed", "theorem"),
("removed_violation_penalty", "coupled", "none"),
):
control_errors: list[float] = []
changed = 0
for graph_index, (risk, labels, ancestors, grid, hard) in enumerate(records):
score, beta, tau_s, lambda_ = soft_score(
risk, labels, ancestors, grid, final_temperature,
beta_mode=beta_mode, violation_mode=violation_mode,
)
error = abs(score - hard)
control_errors.append(error)
changed += int(error > 1e-6)
control_rows.append({
"control": control_name,
"graph_index": graph_index,
"temperature": format(final_temperature, ".17g"),
"hard_score": format(hard, ".17g"),
"control_score": format(score, ".17g"),
"absolute_error": format(error, ".17g"),
"beta": format(beta, ".17g"),
"tau_s": format(tau_s, ".17g"),
"lambda": format(lambda_, ".17g"),
})
control_summaries[control_name] = {
"graphs_changed_beyond_1e-6": changed,
"mean_absolute_error": float(np.mean(control_errors)),
"maximum_absolute_error": max(control_errors),
"all_finite": all(math.isfinite(value) for value in control_errors),
}
final_errors = errors_by_temperature[final_temperature]
initial_errors = errors_by_temperature[TEMPERATURES[0]]
final_quantile_errors = [
float(row["absolute_error"])
for row in quantile_rows
if float(row["temperature"]) == final_temperature
]
gates = {
"official_release_scale_50_graphs_503_nodes": len(records) == 50 and total_nodes == 503,
"theorem_source_markers_exact": all(marker_checks.values()),
"official_dataset_sha256_exact": sha256(data_path) == DATA_SHA256,
"lambda_grid_span_exactly_one_half": all(abs(value - 0.5) <= 1e-15 for value in lambda_spans),
"final_all_50_within_1e-10": sum(error <= FINAL_TOLERANCE for error in final_errors) == 50,
"final_max_error_below_1e-10": max(final_errors) < FINAL_TOLERANCE,
"mean_error_contracts_by_1e8": float(np.mean(initial_errors)) > 1e8 * float(np.mean(final_errors)),
"all_15_final_quantiles_within_1e-10": len(final_quantile_errors) == 15 and max(final_quantile_errors) < FINAL_TOLERANCE,
"quantile_order_statistic_stability_bound": quantile_bound_holds,
"fixed_beta_destructive_control_fails": control_summaries["fixed_beta_8"]["graphs_changed_beyond_1e-6"] == 50 and control_summaries["fixed_beta_8"]["mean_absolute_error"] > 0.1,
"removed_violation_destructive_control_fails": control_summaries["removed_violation_penalty"]["graphs_changed_beyond_1e-6"] >= false_graphs and control_summaries["removed_violation_penalty"]["mean_absolute_error"] > 0.1,
"all_reported_values_finite": all(math.isfinite(value) for values in errors_by_temperature.values() for value in values),
}
if not all(gates.values()):
raise RuntimeError({"failed_gates": [name for name, passed in gates.items() if not passed], "controls": control_summaries})
temperature_summary = []
for temperature in TEMPERATURES:
errors = errors_by_temperature[temperature]
temperature_summary.append({
"temperature": temperature,
"tau_s": temperature ** 0.5,
"beta": temperature ** -1.0,
"maximum_absolute_error": max(errors),
"mean_absolute_error": float(np.mean(errors)),
"graphs_within_1e-6": sum(error <= 1e-6 for error in errors),
"graphs_within_1e-10": sum(error <= FINAL_TOLERANCE for error in errors),
})
summary = {
"status": "pass",
"paper_id": PAPER_ID,
"official_commit": OFFICIAL_COMMIT,
"registered_claim": "Theorem 3.1 calibration convergence and conformal quantile recovery",
"execution_scope": {
"dataset": "released MATH_open_subclaims_with_scores_and_semantic_eval.json",
"graphs": len(records),
"claim_nodes": total_nodes,
"dependency_edges": total_edges,
"graphs_with_false_nodes": false_graphs,
"temperature_schedule": list(TEMPERATURES),
"graph_temperature_evaluations": len(path_rows),
"alpha_grid": list(ALPHAS),
"quantile_evaluations": len(quantile_rows),
},
"theorem_contract": {
"soft_keep": "sigmoid((tau-risk+sqrt(T))/T)",
"tau_s": "T^0.5",
"beta": "T^-1",
"lambda_grid_span": 0.5,
"required_upper_bound": 1.0,
"hard_oracle": "largest threshold whose selected ancestor-coherent subgraph contains no false node",
"soft_sort_shim_used_in_reported_path": False,
},
"source_checks": marker_checks,
"temperature_results": temperature_summary,
"final_temperature_result": temperature_summary[-1],
"quantile_recovery": {
"alphas": len(ALPHAS),
"final_maximum_absolute_error": max(final_quantile_errors),
"order_statistic_stability_bound_holds_all_cells": quantile_bound_holds,
"scope_note": "Quantile recovery is certified as an order-statistic corollary of uniform score convergence; it is not attributed to the theorem text alone.",
},
"destructive_controls": control_summaries,
"gates": gates,
"no_paper_scale_rerun_invented": True,
}
write_csv(args.output_dir / "calibration_limit_path.csv", list(path_rows[0]), path_rows)
write_csv(args.output_dir / "calibration_quantiles.csv", list(quantile_rows[0]), quantile_rows)
write_csv(args.output_dir / "calibration_controls.csv", list(control_rows[0]), control_rows)
write_json(args.output_dir / "calibration_limit_summary.json", summary)
print(json.dumps({
"status": "pass",
"graphs": len(records),
"nodes": total_nodes,
"final_max_error": max(final_errors),
"final_quantile_max_error": max(final_quantile_errors),
"fixed_beta_mean_error": control_summaries["fixed_beta_8"]["mean_absolute_error"],
"removed_violation_mean_error": control_summaries["removed_violation_penalty"]["mean_absolute_error"],
}, sort_keys=True))
if __name__ == "__main__":
main()