#!/usr/bin/env python3
"""Numerical benchmark for the scalar GAD capacity objective in family 276.

The upstream theorem and its selected formal scope have not been independently
verified here. Numeric search and its analytic grid allowance are not an
outward-rounded or formally certified capacity bound.
"""
import argparse
import json
import math
from pathlib import Path

SOURCE_COMMIT = 'fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb'
SOURCE_URL = f'https://github.com/openai/math/blob/{SOURCE_COMMIT}/preprints/Classical-capacity-and-entropy-inequalities-for-generalized-amplitude-damping-September-24-2026/paper.pdf'


def unit_interval(value, name):
    if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value) or not 0 <= value <= 1:
        raise ValueError(f'{name} must be a finite number in [0, 1]')
    return float(value)


def binary_entropy(probability):
    q = unit_interval(probability, 'entropy probability')
    q = min(q, 1 - q)
    if q == 0:
        return 0.0
    return (-q * math.log(q) - (1 - q) * math.log1p(-q)) / math.log(2)


def scalar_holevo(gamma, nu, p):
    """Bits for the equiprobable states sqrt(1-p)|0> ± sqrt(p)|1>."""
    gamma = unit_interval(gamma, 'gamma')
    nu = unit_interval(nu, 'nu')
    p = unit_interval(p, 'signal excited-state population')
    mean_excited = (1 - gamma) * p + gamma * nu
    determinant = gamma * nu * (1 - nu) + gamma * (1 - gamma) * (p - nu) ** 2
    # The exact expression is in [0, 1/4]. Clamp only floating-point edge error.
    if determinant < -1e-14 or determinant > 0.25 + 1e-14:
        raise ArithmeticError('Output determinant outside the state range')
    determinant = min(0.25, max(0.0, determinant))
    root = math.sqrt(max(0.0, 1 - 4 * determinant))
    small_eigenvalue = 2 * determinant / (1 + root)
    return max(0.0, binary_entropy(mean_excited) - binary_entropy(small_eigenvalue))


def _refine(gamma, nu, left, right):
    """Improve a grid bracket; global completeness comes only from the grid."""
    ratio = (math.sqrt(5) - 1) / 2
    x1 = right - ratio * (right - left)
    x2 = left + ratio * (right - left)
    y1 = scalar_holevo(gamma, nu, x1)
    y2 = scalar_holevo(gamma, nu, x2)
    for _ in range(80):
        if y1 < y2:
            left = x1
            x1, y1 = x2, y2
            x2 = left + ratio * (right - left)
            y2 = scalar_holevo(gamma, nu, x2)
        else:
            right = x2
            x2, y2 = x1, y1
            x1 = right - ratio * (right - left)
            y1 = scalar_holevo(gamma, nu, x1)
    return max(((x1, y1), (x2, y2)), key=lambda pair: pair[1])


def _entropy_modulus(delta):
    return 1.0 if delta >= 0.5 else binary_entropy(max(0.0, delta))


def estimate(gamma, nu, grid_intervals=4096):
    gamma = unit_interval(gamma, 'gamma')
    nu = unit_interval(nu, 'nu')
    if isinstance(grid_intervals, bool) or not isinstance(grid_intervals, int) or not 2 <= grid_intervals <= 1000000:
        raise ValueError('grid_intervals must be an integer between 2 and 1000000')
    grid = [scalar_holevo(gamma, nu, i / grid_intervals) for i in range(grid_intervals + 1)]
    best_index = max(range(len(grid)), key=grid.__getitem__)
    p, value = best_index / grid_intervals, grid[best_index]
    # Flat objectives need no repeated local searches. The initial global grid
    # has full finite sampling coverage even if a refinement bracket is multimodal.
    if 0 < best_index < grid_intervals and gamma not in (0.0, 1.0):
        candidate = _refine(gamma, nu, (best_index - 1) / grid_intervals, (best_index + 1) / grid_intervals)
        if candidate[1] > value:
            p, value = candidate
    # Every p is within h=1/(2M) of a grid point. |Δt| <= (1-gamma)h.
    # |Δv| <= 2 gamma(1-gamma) max(nu,1-nu) h. The smaller
    # eigenvalue changes by at most sqrt(|Δv|). Binary entropy's
    # continuity modulus then bounds the two entropy changes in exact math.
    # This numeric allowance does not include floating-point rounding error.
    h = 1 / (2 * grid_intervals)
    delta_mean = (1 - gamma) * h
    delta_eigenvalue = math.sqrt(2 * gamma * (1 - gamma) * max(nu, 1 - nu) * h)
    allowance = _entropy_modulus(delta_mean) + _entropy_modulus(delta_eigenvalue)
    return dict(
        gamma=gamma, stationary_excited_population=nu,
        best_observed_holevo_bits=value, signal_excited_population=p,
        best_grid_holevo_bits=grid[best_index], grid_intervals=grid_intervals,
        analytic_grid_discretization_allowance_bits=allowance,
        estimated_model_maximum_upper_with_grid_allowance_bits=min(1.0, grid[best_index] + allowance),
        ensemble='equiprobable sqrt(1-p)|0> plus/minus sqrt(p)|1>',
        decoding='asymptotic collective decoding; separate-output decoding not certified',
        evidence='floating_point_model_benchmark_not_formal_or_outward_rounded',
        rounding_error_bound=None, independent_theorem_verification='not_run',
        source_commit=SOURCE_COMMIT, source_url=SOURCE_URL,
    )


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--gamma', type=float, required=True)
    parser.add_argument('--nu', type=float, required=True, help='Stationary excited-state population, not ground-state population')
    parser.add_argument('--grid-intervals', type=int, default=4096)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    result = estimate(args.gamma, args.nu, args.grid_intervals)
    body = json.dumps(result, indent=2, allow_nan=False) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(body)
    else:
        print(body, end='')


if __name__ == '__main__':
    main()
