Download calibration_limit_audit.py from SabaPivot/repro-differentiable-conformal-training-for-llm-reasoning-factuality: direct link, hf CLI and curl.
- Browser
- Download file 17.4 kB
-
https://huggingface.co/spaces/SabaPivot/repro-differentiable-conformal-training-for-llm-reasoning-factuality/resolve/main/calibration_limit_audit.py
- Command line
-
hf download hf://spaces/SabaPivot/repro-differentiable-conformal-training-for-llm-reasoning-factuality/calibration_limit_audit.py
-
curl -L -o calibration_limit_audit.py https://huggingface.co/spaces/SabaPivot/repro-differentiable-conformal-training-for-llm-reasoning-factuality/resolve/main/calibration_limit_audit.py
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() | |