meteora-repro / code /theorem6_lower_bound.py
ProCreations's picture
Audit Theorem 6 consistency robustness lower bound
f601267
Raw
History Blame Contribute Delete
5.22 kB
#!/usr/bin/env python3
"""CPU-only finite audit of the paper's Theorem 6 lower-bound construction.
Decimal arithmetic keeps the tail probabilities and policy-cost comparisons
well above double precision. The construction is the paper's two-phase
prediction tail; the adversarial truth is a point at b**100, so its support is
finite for every reported b without materializing that enormous integer.
"""
from __future__ import annotations
import argparse
import json
import math
from decimal import Decimal, getcontext
from pathlib import Path
getcontext().prec = 70
D = Decimal
def phase_parameters(b: int) -> tuple[Decimal, Decimal, int]:
bd = D(b)
r = D(1) - D(2) / bd
slow = D(1) - D(1) / (D(2) * bd)
K = math.ceil(0.5 * b * math.log(b))
return r, slow, K
def tail(b: int, t: int) -> Decimal:
r, slow, K = phase_parameters(b)
if t <= K:
return r**t
return (r**K) * (slow ** (t - K))
def policy_cost(b: int, B: int) -> Decimal:
"""Expected cost of buying after B rental days under the prediction tail."""
r, slow, K = phase_parameters(b)
bd = D(b)
if B <= K:
rent_sum = (D(1) - r**B) / (D(1) - r)
else:
first = (D(1) - r**K) / (D(1) - r)
second = (r**K) * (D(1) - slow ** (B - K)) / (D(1) - slow)
rent_sum = first + second
return rent_sum + bd * tail(b, B)
def theorem6_rows() -> list[dict]:
rows = []
for b in (2**8, 2**10, 2**12, 2**14, 2**16, 2**18, 2**20):
r, slow, K = phase_parameters(b)
bd = D(b)
sqrt_b = bd.sqrt()
log_b = bd.ln()
qK = r**K
def excess(B: int) -> Decimal:
return (bd / D(2)) * (r**B - qK)
constants = []
for C in (D("0.5"), D("1"), D("2")):
target = D(2) * C * sqrt_b
lo, hi = 0, K
while lo < hi:
mid = (lo + hi + 1) // 2
if excess(mid) >= target:
lo = mid
else:
hi = mid - 1
B0 = lo
constants.append(
{
"C": float(C),
"B0": B0,
"B0_over_b_log_b": float(D(B0) / (bd * log_b)),
"randomized_lower_bound_EB_over_b_log_b": float((D(B0) / D(2)) / (bd * log_b)),
"delta_B0_over_sqrt_b": float(excess(B0) / sqrt_b),
"delta_B0_plus_1_over_sqrt_b": float(excess(B0 + 1) / sqrt_b),
}
)
N = K + math.ceil(4.0 * b * math.log(b))
# The finite-support replacement can aggregate this remaining tail at N.
rows.append(
{
"b": b,
"K": K,
"finite_support_cutoff_N": N,
"tail_at_N": float(tail(b, N)),
"prediction_optimal_threshold": K,
"constants": constants,
}
)
return rows
def fresh_protection_rows() -> dict:
# New scales, distinct from the original b=2^4,...,2^20 every-two-exponents grid.
scales = (2**7, 2**9, 2**11, 2**13, 2**15, 2**17, 2**19)
identity_errors = []
consistency = []
robustness = []
thresholds = []
for b in scales:
r, slow, K = phase_parameters(b)
sqrt_b = math.isqrt(b)
u = 0
while tail(b, u) > D(1) / D(sqrt_b):
u += 1
chosen = u + sqrt_b
thresholds.append(chosen)
# Exact Decimal recurrence identity across the two phases and the clamp point.
for B in (0, 1, max(0, u - 1), u, K, K + sqrt_b):
lhs = policy_cost(b, B + 1) - policy_cost(b, B)
rhs = tail(b, B) - D(b) * (tail(b, B) - tail(b, B + 1))
identity_errors.append(abs(lhs - rhs))
optimum = policy_cost(b, K)
selected = policy_cost(b, chosen)
consistency.append(float((selected - optimum) / D(sqrt_b)))
robustness.append(float(D(chosen) / (D(b) * D(b).ln())))
return {
"scales": list(scales),
"policy_cost_identity_cells": len(identity_errors),
"max_policy_cost_identity_error": float(max(identity_errors)),
"consistency_loss_over_sqrt_b": consistency,
"robustness_loss_over_b_log_b": robustness,
"thresholds_for_five_truths_each": [thresholds[i] for i in range(len(thresholds)) for _ in range(5)],
"distinct_thresholds_per_prediction": 1,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
rows = theorem6_rows()
result = {
"theorem6": rows,
"fresh_protection": fresh_protection_rows(),
"randomized_argument": "For every C, B0 is the largest threshold with Delta(B0)>=2*C*sqrt(b); any threshold mixture with E[Delta]<=C*sqrt(b) has Pr(B<=B0)<=1/2 and therefore E[B]>=B0/2.",
"adversarial_truth": "point mass at T=b**100; OPT buys immediately and additive loss equals E[B] for every finite threshold mixture",
}
args.output.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8")
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()