#!/usr/bin/env python3
"""用精确分数复算早停风险包络与有限网格选择。

运行：python3 early-stopping-grid-reader.py
只用 Python 标准库；输出 JSON，包括两个完整网格与边界检查。
所有小数输入应写成 Fraction 或分数字符串，不接受浮点近似。

本程序逐网格点独立计算整数幂，不采用正文也介绍的逐轮扫描。
设特征值输入数为 q（可含零）、网格列表长度为 m、最大轮数为 M，
算术/比较次数为 O(m*q*log(2+M) + m*log(1+m))。这不是位复杂度界；
精确分数的分子分母长度随轮数增长。

choose_stop 不接收标签或真信号。风险证书要求：固定设计线性真模型、
可靠参数半径、均值零且精确各向同性协方差、设计先定的步长与网格。
以下有限计算不能代替正文对全部参数球的证明，也不验证这些统计假设。
"""
from fractions import Fraction as F
import json


def rational(x):
    if isinstance(x, bool) or not isinstance(x, (int, str, F)):
        raise TypeError("Use integers, Fraction values, or exact fraction strings")
    return F(x)


def nonnegative_integer(x, name):
    if isinstance(x, bool) or not isinstance(x, int) or x < 0:
        raise ValueError(f"{name} must be a nonnegative integer")
    return x


def model(eigenvalues, eta, radius_squared, noise_variance, n):
    d = tuple(rational(x) for x in eigenvalues)
    eta, radius_squared, noise_variance = map(
        rational, (eta, radius_squared, noise_variance)
    )
    nonnegative_integer(n, "n")
    if n == 0 or eta <= 0 or radius_squared < 0 or noise_variance < 0:
        raise ValueError("Need n>=1, eta>0, and nonnegative squared budgets")
    if any(x < 0 for x in d) or sum(x > 0 for x in d) > n:
        raise ValueError("A design spectrum needs nonnegative values and rank<=n")
    if any(eta * x > 1 for x in d):
        raise ValueError("Need eta*d<=1 in every direction")
    return d, eta, radius_squared, noise_variance, n


def risk_envelope(eigenvalues, eta, radius_squared, noise_variance, n, t):
    """返回最坏偏差平方、精确方差、包络；零特征值可保留或省略。"""
    d, eta, r2, s2, n = model(eigenvalues, eta, radius_squared, noise_variance, n)
    nonnegative_integer(t, "t")
    # 显式保留零轮语义，不依赖实现对 0**0 的选择。
    a = tuple(F(1) if t == 0 else (1 - eta * x) ** t for x in d)
    weights = tuple(x * h * h for x, h in zip(d, a))
    bias = r2 * max(weights, default=F(0))
    variance = s2 / n * sum(((1 - h) ** 2 for h in a), F(0))
    return {"t": t, "worst_case_bias": bias, "variance": variance,
            "risk_upper": bias + variance}


def choose_stop(eigenvalues, eta, radius_squared, noise_variance, n, grid):
    """事前给定网格中取最小包络，并列时选最早轮次；不读取标签。"""
    grid = tuple(grid)
    if not grid:
        raise ValueError("The declared grid must be nonempty")
    for t in grid:
        nonnegative_integer(t, "grid element")
    d = tuple(eigenvalues)
    rows = [risk_envelope(d, eta, radius_squared, noise_variance, n, t)
            for t in sorted(set(grid))]
    best = min(rows, key=lambda row: (row["risk_upper"], row["t"]))
    return {"best": best, "rows": rows}


def fixed_signal_risk(eigenvalues, eta, noise_variance, n, t, coefficient_squares):
    """仅供复算具体信号；网格选择不调用此函数，不用信号信息调参。"""
    d = tuple(rational(x) for x in eigenvalues)
    b2 = tuple(rational(x) for x in coefficient_squares)
    if len(d) != len(b2) or any(x < 0 for x in b2):
        raise ValueError("Need one nonnegative coefficient square per eigenvalue")
    row = risk_envelope(d, eta, sum(b2), noise_variance, n, t)
    a = tuple(F(1) if t == 0 else (1 - rational(eta) * x) ** t for x in d)
    return sum((x * h * h * b for x, h, b in zip(d, a, b2)), F(0)) + row["variance"]


def require(condition, message):
    # 使用异常而不是 assert，python -O 下也保持检查。
    if not condition:
        raise RuntimeError(message)


def check_examples():
    d, eta, r2, s2, n = (F(1), F(1, 10)), F(1, 2), F(2), F(1), 2
    grid = tuple(range(13))
    base = choose_stop(d, eta, r2, s2, n, grid)
    best = base["best"]
    require(best["t"] == 2, "base stopping choice")
    require(best["worst_case_bias"] == F(130321, 800000), "base bias")
    require(best["variance"] == F(91521, 320000), "base variance")
    require(best["risk_upper"] == F(718247, 1600000), "base envelope")
    require(base["rows"][4]["variance"] > F(9, 20) > best["risk_upper"],
            "variance excludes rounds 4..12")
    actual = fixed_signal_risk(d, eta, s2, n, 2, (2, 0))
    attained = fixed_signal_risk(d, eta, s2, n, 2, (0, 2))
    require(actual == F(131521, 320000) < attained == best["risk_upper"],
            "specific signal and attaining signal")
    # 对每轮分别把全部半径分给最大偏差方向，核对可达性。
    for row in base["rows"]:
        t = row["t"]
        j = max(range(len(d)), key=lambda k: d[k] * (1 - eta * d[k]) ** (2 * t))
        b2 = tuple(r2 if k == j else F(0) for k in range(len(d)))
        require(fixed_signal_risk(d, eta, s2, n, t, b2) == row["risk_upper"],
                f"attaining direction at t={t}")

    transfer = choose_stop(d, 1, r2, s2, n, grid)
    require(transfer["best"]["t"] == 3, "changed-step transfer choice")
    require(transfer["best"]["risk_upper"] == F(6430087, 10000000), "transfer risk")
    require(transfer["rows"][4]["risk_upper"] == F(645227047, 1000000000),
            "transfer round 4")
    require(transfer["rows"][0]["risk_upper"] == 2, "zero round at eta*d=1")
    for row in transfer["rows"][1:]:
        x = F(9, 10) ** row["t"]
        require(row["risk_upper"] == 1 - x + F(7, 10) * x * x,
                "transfer quadratic")
    with_zero = choose_stop((*d, F(0)), 1, r2, s2, n, grid)
    require(with_zero == transfer, "adding a zero feature preserves prediction risk")
    require(choose_stop((0, 0), 1, r2, s2, n, (9, 3, 6))["best"]["t"] == 3,
            "zero design tie rule")
    require(choose_stop((), 1, r2, s2, n, (0, 1))["best"]["risk_upper"] == 0,
            "empty positive spectrum")
    require(choose_stop(d, eta, 0, s2, n, grid)["best"]["t"] == 0,
            "zero radius")
    require(choose_stop(d, eta, r2, 0, n, grid)["best"]["t"] == 12,
            "zero noise")

    bad_r = choose_stop(d, eta, 1, s2, n, grid)
    require(bad_r["best"]["t"] == 2 and
            bad_r["best"]["risk_upper"] == F(293963, 800000) < attained,
            "understated radius violates advertised bound")
    real_four_noise = fixed_signal_risk(d, eta, 4, n, 2, (0, 2))
    require(real_four_noise == F(1045531, 800000) > best["risk_upper"],
            "understated noise violates advertised bound")
    four_noise = choose_stop(d, eta, r2, 4, n, grid)
    require(four_noise["best"]["t"] == 1 and
            four_noise["best"]["risk_upper"] == F(201, 200),
            "correct larger noise changes the choice")
    # 零信号、零噪声也允许正的次高斯上尺度；不能由上尺度推出等号。
    zero_actual = fixed_signal_risk(d, eta, 0, n, 1, (0, 0))
    proxy_bound = risk_envelope(d, eta, 0, 1, n, 1)["risk_upper"]
    require(zero_actual == 0 < proxy_bound, "upper proxy is not an exact variance")

    invalid = [(d, eta, r2, s2, n, ()), (d, 2, r2, s2, n, grid),
               (d, eta, -1, s2, n, grid), (d, eta, r2, -1, n, grid),
               (d, eta, r2, s2, 0, grid), (d, eta, r2, s2, 1, grid),
               (d, eta, r2, s2, n, (-1,)), ((-1,), eta, r2, s2, n, grid)]
    for args in invalid:
        try:
            choose_stop(*args)
        except (ValueError, TypeError):
            pass
        else:
            raise RuntimeError("invalid input was accepted")
    return {
        "contract": "fixed design; exact isotropic noise covariance; known radius; response-free finite grid",
        "base_grid": base,
        "step_one_transfer_grid": transfer,
        "checks": {
            "passed": True,
            "old_specific_signal_risk": actual,
            "attaining_signal_risk": attained,
            "understated_radius_bound": bad_r["best"]["risk_upper"],
            "actual_noise_four_risk_at_two": real_four_noise,
            "correct_noise_four_choice": four_noise["best"],
            "zero_noise_actual_risk": zero_actual,
            "positive_proxy_bound": proxy_bound,
            "zero_round_zero_spectrum_zero_radius_zero_noise": "passed",
            "invalid_inputs": len(invalid),
        },
    }


if __name__ == "__main__":
    print(json.dumps(check_examples(), default=str, ensure_ascii=False, indent=2))
