Risk-Calibrated World-Model Gates / risk_gate_experiment.py
Failed on benchmark
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()