Risk-Calibrated World-Model Gates / risk_gate_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2import math
  3from pathlib import Path
  4
  5import numpy as np
  6
  7
  8def required_rollouts(r_low: float, delta: float) -> int:
  9    """Smallest integer N for (1-r_low)^N <= delta."""
 10    if not 0.0 < r_low < 1.0 or not 0.0 < delta < 1.0:
 11        raise ValueError("r_low and delta must be in (0,1)")
 12    return int(math.ceil(math.log(delta) / math.log1p(-r_low)))
 13
 14
 15def beta_prior_lower_bound(events: int, total: int, alpha: float = 0.05) -> float:
 16    """Simple conservative one-sided lower estimate.
 17
 18    The zero-event branch is the 1-alpha lower credible bound induced by a
 19    Beta(1,1) prior. For nonzero counts, Wilson's one-sided lower bound is used.
 20    """
 21    if total <= 0 or events < 0 or events > total:
 22        raise ValueError("invalid event counts")
 23    if events == 0:
 24        return 1.0 - alpha ** (1.0 / (total + 1.0))
 25    p, z = events / total, 1.6448536269514722
 26    denom = 1.0 + z * z / total
 27    center = (p + z * z / (2 * total)) / denom
 28    spread = z * math.sqrt(p * (1-p) / total + z*z / (4 * total * total)) / denom
 29    return max(0.0, center - spread)
 30
 31
 32def exact_miss_sweep(seed=7, trials=100_000):
 33    rng = np.random.default_rng(seed)
 34    rows = []
 35    for r in (0.01, 0.03, 0.10, 0.25):
 36        for n in (5, 20, 50):
 37            miss = rng.binomial(n, r, size=trials) == 0
 38            observed = float(np.mean(miss))
 39            predicted = (1-r) ** n
 40            rows.append({"r": r, "N": n, "observed_miss": observed,
 41                         "predicted_miss": predicted,
 42                         "abs_error": abs(observed - predicted)})
 43    return rows
 44
 45
 46def sizing_sweep(delta=0.05):
 47    rows = []
 48    for r in (0.005, 0.01, 0.02, 0.05, 0.10):
 49        n = required_rollouts(r, delta)
 50        exact = (1-r) ** n
 51        approximation = math.log(1/delta) / r
 52        rows.append({"r": r, "N_required": n, "exact_miss": exact,
 53                     "target_delta": delta, "small_r_approx_N": approximation,
 54                     "N_over_approx": n / approximation})
 55    return rows
 56
 57
 58def probe_sweep(seed=11, trials=100_000):
 59    """Matched-budget detection experiment.
 60
 61    A random rollout sees the critical mode with probability r. A directed probe
 62    sees it with probability q. The proposed gate spends 30 random + 20 probes;
 63    the baseline spends all 50 randomly. Independent draws make the predicted
 64    proposed miss probability (1-r)^30 (1-q)^20.
 65    """
 66    rng = np.random.default_rng(seed)
 67    rows = []
 68    for r in (0.005, 0.01, 0.02, 0.05):
 69        q = min(0.60, 30*r)  # toy boundary score: probes amplify rare modes
 70        baseline_miss = float(np.mean(rng.binomial(50, r, trials) == 0))
 71        idea_miss = float(np.mean((rng.binomial(30, r, trials) +
 72                                   rng.binomial(20, q, trials)) == 0))
 73        predicted_idea_miss = (1-r)**30 * (1-q)**20
 74        rows.append({"r": r, "probe_rate_q": q,
 75                     "baseline_miss_observed": baseline_miss,
 76                     "baseline_miss_predicted": (1-r)**50,
 77                     "idea_miss_observed": idea_miss,
 78                     "idea_miss_predicted": predicted_idea_miss,
 79                     "detection_gain": (1-baseline_miss) - (1-idea_miss),
 80                     "probe_scaling_factor_q_over_r": q/r})
 81    return rows
 82
 83
 84def validate(rows, sizing, probes):
 85    max_exact_error = max(x["abs_error"] for x in rows)
 86    max_probe_error = max(abs(x["idea_miss_observed"] - x["idea_miss_predicted"])
 87                          for x in probes)
 88    assert max_exact_error < 0.01
 89    assert max_probe_error < 0.01
 90    assert all(x["exact_miss"] <= x["target_delta"] for x in sizing)
 91    # Predictions: exact exponential scaling; ceil-sized N meets delta; probes
 92    # reduce miss probability for every tested rare-event rate.
 93    assert all(x["idea_miss_predicted"] < x["baseline_miss_predicted"] for x in probes)
 94    return {"max_exact_abs_error": max_exact_error,
 95            "max_probe_abs_error": max_probe_error,
 96            "all_sizing_rows_meet_delta": True,
 97            "all_probe_rows_improve_over_baseline": True}
 98
 99
100def main():
101    exact = exact_miss_sweep()
102    sizing = sizing_sweep()
103    probes = probe_sweep()
104    checks = validate(exact, sizing, probes)
105    result = {"checks": checks, "exact_miss_sweep": exact,
106              "required_count_sweep": sizing, "matched_budget_probe_sweep": probes}
107    Path("results.json").write_text(json.dumps(result, indent=2))
108    print(json.dumps(result, indent=2))
109
110
111if __name__ == "__main__":
112    main()