generate.py

back to table · edit · history · where entries came from · files · download

19534 bytes, as of the version from 2026-09-25 22:17 (current). Recorded here, not run.

"""Characteristic values a_n(q) of the Mathieu equation for NumberDB T436.

This generator computes the characteristic values in the DLMF convention

    w'' + (a - 2 q cos(2 z)) w = 0

for the existing T436 argument set.  It is designed to run offline and to
record the convergence checks used for the precision repair.

Run it with SageMath:

    $ sage -pip install numberdb          # once
    $ sage -python generate.py --write-computation computation.json
    $ sage -python generate.py --self-check

Publishing is not part of this repair job.  The optional --publish path requires
an explicit --allow-publish flag and is not run by the precision queue.
"""

import argparse
import json
import math
import os
import sys
from functools import lru_cache

import mpmath as mp
import numberdb.sage as numberdb
from numberdb._write import to_text
from sage.rings.integer_ring import ZZ
from sage.rings.rational_field import QQ
from sage.rings.real_mpfi import RealIntervalField
from sage.rings.real_mpfr import RealField


TABLE = "T436"
DIGITS = 100
MAX_N = 10
Q_TEXTS = (
    "1/8", "1/6", "1/4", "1/3", "1/2", "2/3", "3/4", "1",
    "4/3", "3/2", "2", "3", "4", "6", "8",
)
Q_VALUES = tuple(QQ(text) for text in Q_TEXTS)

# Candidate digits are written from the agreement interval produced by these
# four finite-section computations.
GENERATION_TRUNCATIONS = (90, 120)
GENERATION_WORKING_DIGITS = (220, 260)

# Every stored entry is checked against this stronger finite-section run.
CHECK_TRUNCATION = 150
CHECK_WORKING_DIGITS = 320

BISECTION_GUARD_DIGITS = 28
BIT_GUARD = 96
ODE_CHECK_CASES = (("1/8", 0), ("1", 4), ("8", 0), ("8", 1), ("8", 10))


def bits(decimal_digits, guard=BIT_GUARD):
    return numberdb.bits(decimal_digits, losing=guard)


def real(q, field):
    return field(QQ(q))


def a_even_tridiagonal(q, size, prec):
    field = RealField(prec)
    q = real(q, field)
    diagonal = tuple(field(2 * r) ** 2 for r in range(size))
    off_diagonal = [q] * (size - 1)
    if off_diagonal:
        off_diagonal[0] = field(2).sqrt() * q
    return field, diagonal, tuple(off_diagonal)


def a_odd_tridiagonal(q, size, prec):
    field = RealField(prec)
    q = real(q, field)
    diagonal = [field(2 * r + 1) ** 2 for r in range(size)]
    diagonal[0] += q
    off_diagonal = tuple(q for _ in range(size - 1))
    return field, tuple(diagonal), off_diagonal


def b_even_tridiagonal(q, size, prec):
    field = RealField(prec)
    q = real(q, field)
    diagonal = tuple(field(2 * (r + 1)) ** 2 for r in range(size))
    off_diagonal = tuple(q for _ in range(size - 1))
    return field, diagonal, off_diagonal


def b_odd_tridiagonal(q, size, prec):
    field = RealField(prec)
    q = real(q, field)
    diagonal = [field(2 * r + 1) ** 2 for r in range(size)]
    diagonal[0] -= q
    off_diagonal = tuple(q for _ in range(size - 1))
    return field, tuple(diagonal), off_diagonal


def tridiagonal(kind, q, size, prec):
    if kind == "a-even":
        return a_even_tridiagonal(q, size, prec)
    if kind == "a-odd":
        return a_odd_tridiagonal(q, size, prec)
    if kind == "b-even":
        return b_even_tridiagonal(q, size, prec)
    if kind == "b-odd":
        return b_odd_tridiagonal(q, size, prec)
    raise ValueError("unknown matrix kind %r" % (kind,))


def sturm_count_less(diagonal, off_diagonal, x):
    field = x.parent()
    tiny = field(2) ** (-(field.precision() - 8))
    negative = 0
    previous = None
    for index, diagonal_entry in enumerate(diagonal):
        pivot = diagonal_entry - x
        if index:
            if previous == 0:
                previous = tiny
            pivot -= off_diagonal[index - 1] ** 2 / previous
        if pivot < 0:
            negative += 1
        if pivot == 0:
            pivot = -tiny
        previous = pivot
    return negative


def bisect_eigenvalue(diagonal, off_diagonal, index, digits):
    field = diagonal[0].parent()
    radius = sum(abs(entry) for entry in off_diagonal) + field(1)
    low = min(diagonal) - radius
    high = max(diagonal) + radius
    if sturm_count_less(diagonal, off_diagonal, low) > index:
        raise ArithmeticError("left bracket is too high for eigenvalue %s" % index)
    if sturm_count_less(diagonal, off_diagonal, high) <= index:
        raise ArithmeticError("right bracket is too low for eigenvalue %s" % index)
    tolerance = field(10) ** (-(digits + BISECTION_GUARD_DIGITS))
    while high - low > tolerance:
        middle = (low + high) / field(2)
        if sturm_count_less(diagonal, off_diagonal, middle) <= index:
            low = middle
        else:
            high = middle
    return (low + high) / field(2)


def spectrum_length(kind):
    return MAX_N // 2 + 1


@lru_cache(maxsize=None)
def spectrum(kind, q_text, size, decimal_digits, digits):
    prec = bits(decimal_digits)
    _field, diagonal, off_diagonal = tridiagonal(kind, QQ(q_text), size, prec)
    return tuple(
        bisect_eigenvalue(diagonal, off_diagonal, index, digits)
        for index in range(spectrum_length(kind))
    )


def a_value(q, n, size, decimal_digits, digits=DIGITS):
    kind = "a-even" if n % 2 == 0 else "a-odd"
    return spectrum(kind, str(QQ(q)), size, decimal_digits, digits)[n // 2]


def b_value(q, n, size, decimal_digits, digits=DIGITS):
    if n < 1:
        raise ValueError("b_n starts at n=1")
    kind = "b-even" if n % 2 == 0 else "b-odd"
    index = n // 2 - 1 if n % 2 == 0 else (n - 1) // 2
    return spectrum(kind, str(QQ(q)), size, decimal_digits, digits)[index]


def generation_values(q, n):
    return [
        a_value(q, n, size, working, DIGITS)
        for size in GENERATION_TRUNCATIONS
        for working in GENERATION_WORKING_DIGITS
    ]


def interval_from_values(values):
    field = RealIntervalField(bits(CHECK_WORKING_DIGITS + 20))
    interval = field(values[0])
    for value in values[1:]:
        interval = interval.union(field(value))
    return interval


def value_text(q, n):
    interval = interval_from_values(generation_values(q, n))
    text = to_text(interval, DIGITS, None)
    if text.startswith("["):
        raise ArithmeticError(
            "agreement interval too wide for %s digits at q=%s, n=%s: %s"
            % (DIGITS, q, n, text)
        )
    return text


def compute_entries():
    return {
        str(q): {str(n): value_text(q, n) for n in range(MAX_N + 1)}
        for q in Q_VALUES
    }


def center_of_text(text, decimal_digits=160):
    return RealField(bits(decimal_digits))(text)


def relative_difference(x, y):
    field = RealField(bits(CHECK_WORKING_DIGITS + 20))
    x = field(x)
    y = field(y)
    scale = max(abs(x), abs(y), field(1))
    return abs(x - y) / scale


def decimal_digits_from_relative(relative):
    if relative == 0:
        return 999
    value = float(relative)
    if value == 0.0:
        return 999
    return max(0, int(math.floor(-math.log10(value))))


def small_q_reference(n, q):
    q = QQ(q)
    if n == 0:
        return (-QQ(1) / QQ(2) * q ** 2
                + QQ(7) / QQ(128) * q ** 4
                - QQ(29) / QQ(2304) * q ** 6
                + QQ(68687) / QQ(18874368) * q ** 8)
    if n == 1:
        return (QQ(1) + q
                - QQ(1) / QQ(8) * q ** 2
                - QQ(1) / QQ(64) * q ** 3
                - QQ(1) / QQ(1536) * q ** 4
                + QQ(11) / QQ(36864) * q ** 5
                + QQ(49) / QQ(589824) * q ** 6
                + QQ(55) / QQ(9437184) * q ** 7
                - QQ(83) / QQ(35389440) * q ** 8)
    if n == 2:
        return (QQ(4)
                + QQ(5) / QQ(12) * q ** 2
                - QQ(763) / QQ(13824) * q ** 4
                + QQ(1002401) / QQ(79626240) * q ** 6
                - QQ(1669068401) / QQ(458647142400) * q ** 8)
    if n == 3:
        return (QQ(9)
                + QQ(1) / QQ(16) * q ** 2
                + QQ(1) / QQ(64) * q ** 3
                + QQ(13) / QQ(20480) * q ** 4
                - QQ(5) / QQ(16384) * q ** 5
                - QQ(1961) / QQ(23592960) * q ** 6
                - QQ(609) / QQ(104857600) * q ** 7)
    if n == 4:
        return (QQ(16)
                + QQ(1) / QQ(30) * q ** 2
                + QQ(433) / QQ(864000) * q ** 4
                - QQ(5701) / QQ(2721600000) * q ** 6)
    if n == 5:
        return (QQ(25)
                + QQ(1) / QQ(48) * q ** 2
                + QQ(11) / QQ(774144) * q ** 4
                + QQ(1) / QQ(147456) * q ** 5
                + QQ(37) / QQ(891813888) * q ** 6)
    if n == 6:
        return (QQ(36)
                + QQ(1) / QQ(70) * q ** 2
                + QQ(187) / QQ(43904000) * q ** 4
                + QQ(6743617) / QQ(92935987200000) * q ** 6)
    raise ValueError("no small-q reference for n=%s" % n)


def ode_shooting_residual(q, n, value, dps=130):
    mp.mp.dps = dps
    q = QQ(q)
    qmp = mp.mpf(str(q.numerator())) / mp.mpf(str(q.denominator()))
    amp = mp.mpf(value.str(digits=dps + 10, no_sci=False))

    def rhs(z, y):
        return [y[1], -(amp - 2 * qmp * mp.cos(2 * z)) * y[0]]

    tol = mp.mpf(10) ** (-(dps - 20))
    solution = mp.odefun(rhs, mp.mpf("0"), (mp.mpf("1"), mp.mpf("0")), tol=tol)
    endpoint = solution(mp.pi / 2)
    return endpoint[1] if n % 2 == 0 else endpoint[0]


def run_controls(entries):
    field = RealField(bits(CHECK_WORKING_DIGITS + 20))
    controls = []

    q_zero_max = field(0)
    for n in range(MAX_N + 1):
        value = a_value(QQ(0), n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
        q_zero_max = max(q_zero_max, abs(field(value) - field(ZZ(n) ** 2)))
    if q_zero_max > field("1e-125"):
        raise ArithmeticError("q=0 control failed: %s" % q_zero_max)
    controls.append({
        "name": "q=0 exact limit",
        "scope": "n=0..10",
        "max_absolute_residual": q_zero_max.str(digits=12),
    })

    negative_q_max = field(0)
    for q in Q_VALUES:
        for n in range(0, MAX_N + 1, 2):
            plus = a_value(q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
            minus = a_value(-q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
            negative_q_max = max(negative_q_max, abs(field(plus) - field(minus)))
    if negative_q_max > field("1e-120"):
        raise ArithmeticError("negative-q even symmetry failed: %s" % negative_q_max)
    controls.append({
        "name": "DLMF symmetry a_2m(-q)=a_2m(q)",
        "scope": "all stored q and even n",
        "max_absolute_residual": negative_q_max.str(digits=12),
    })

    interlacing_min_gap = None
    for q in Q_VALUES:
        chain = []
        for n in range(MAX_N + 1):
            chain.append(("a", n, a_value(q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)))
            chain.append(("b", n + 1, b_value(q, n + 1, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)))
        for left, right in zip(chain, chain[1:]):
            gap = field(right[2]) - field(left[2])
            if gap <= 0:
                raise ArithmeticError(
                    "interlacing failed at q=%s between %s_%s and %s_%s"
                    % (q, left[0], left[1], right[0], right[1])
                )
            interlacing_min_gap = gap if interlacing_min_gap is None else min(interlacing_min_gap, gap)
    controls.append({
        "name": "interlacing with companion b_n(q)",
        "scope": "all stored q, a_0 through a_10 and b_1 through b_11",
        "minimum_gap": interlacing_min_gap.str(digits=12),
    })

    q = QQ(1) / QQ(10)
    small_q_max = field(0)
    for n in range(7):
        mine = a_value(q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
        reference = field(small_q_reference(n, q))
        small_q_max = max(small_q_max, abs(field(mine) - reference))
    if small_q_max > field("1e-12"):
        raise ArithmeticError("small-q expansion control failed: %s" % small_q_max)
    controls.append({
        "name": "DLMF small-q expansions through the stated terms",
        "scope": "auxiliary q=1/10, n=0..6",
        "max_absolute_residual": small_q_max.str(digits=12),
        "note": "This is a convention/source check; the truncated DLMF series is not a 100-digit reference.",
    })

    try:
        from scipy import special
    except Exception as exc:
        controls.append({
            "name": "SciPy double-precision comparison",
            "scope": "skipped",
            "note": "SciPy import failed: %r" % (exc,),
        })
    else:
        scipy_max = field(0)
        scipy_worst = None
        for q in Q_VALUES:
            for n in range(MAX_N + 1):
                mine = center_of_text(entries[str(q)][str(n)], 80)
                theirs = field(str(special.mathieu_a(n, float(q))))
                rel = relative_difference(mine, theirs)
                if rel > scipy_max:
                    scipy_max = rel
                    scipy_worst = {"q": str(q), "n": n}
        if scipy_max > field("5e-12"):
            raise ArithmeticError("SciPy comparison failed: %s at %s" % (scipy_max, scipy_worst))
        controls.append({
            "name": "SciPy mathieu_a double-precision comparison",
            "scope": "all stored entries",
            "max_relative_difference": scipy_max.str(digits=12),
            "worst_case": scipy_worst,
        })

    ode_results = []
    for q_text, n in ODE_CHECK_CASES:
        q = QQ(q_text)
        value = a_value(q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
        residual = ode_shooting_residual(q, n, value, dps=130)
        ode_results.append({
            "q": q_text,
            "n": n,
            "residual": mp.nstr(residual, 25),
        })
    controls.append({
        "name": "independent ODE shooting residuals",
        "scope": "selected ordinary and endpoint cases on [0, pi/2]",
        "cases": ode_results,
        "note": "Even n uses y'(pi/2), odd n uses y(pi/2), with y(0)=1 and y'(0)=0.",
    })

    return controls


def check_entries(entries):
    field = RealField(bits(CHECK_WORKING_DIGITS + 20))
    worst_generation_spread = field(0)
    worst_generation_case = None
    worst_check_relative = field(0)
    worst_check_case = None
    min_supported_digits = 999

    for q in Q_VALUES:
        for n in range(MAX_N + 1):
            generated = generation_values(q, n)
            spread = max(field(value) for value in generated) - min(field(value) for value in generated)
            if spread > worst_generation_spread:
                worst_generation_spread = spread
                worst_generation_case = {"q": str(q), "n": n}

            check = a_value(q, n, CHECK_TRUNCATION, CHECK_WORKING_DIGITS, DIGITS)
            stored = center_of_text(entries[str(q)][str(n)], 160)
            relative = relative_difference(stored, check)
            supported = decimal_digits_from_relative(relative)
            min_supported_digits = min(min_supported_digits, supported)
            if relative > worst_check_relative:
                worst_check_relative = relative
                worst_check_case = {"q": str(q), "n": n}

    return {
        "all_entries_checked": True,
        "entries": len(Q_VALUES) * (MAX_N + 1),
        "worst_generation_spread": worst_generation_spread.str(digits=12),
        "worst_generation_case": worst_generation_case,
        "worst_check_relative_difference": worst_check_relative.str(digits=12),
        "worst_check_case": worst_check_case,
        "minimum_supported_digits_from_stronger_finite_section": min_supported_digits,
    }


class MathieuCharacteristicValuesA(numberdb.Generator):
    table = TABLE
    parameters = ("q", "n")
    type = "R"
    digits = DIGITS
    rigour = "heuristic (agreement-checked)"

    def enumerate(self):
        for q in Q_VALUES:
            for n in range(MAX_N + 1):
                yield {"q": str(q), "n": str(n)}

    def value(self, params, digits):
        if digits != DIGITS:
            raise ValueError("this repair generator is calibrated for %s digits" % DIGITS)
        q = QQ(params["q"])
        n = int(params["n"])
        if q not in Q_VALUES or not (0 <= n <= MAX_N):
            raise ValueError("entry outside the T436 repair range: q=%s, n=%s" % (q, n))
        return value_text(q, n)


def write_computation(path):
    entries = compute_entries()
    metrics = check_entries(entries)
    controls = run_controls(entries)
    payload = {
        "table": TABLE,
        "digits": DIGITS,
        "q_values": Q_TEXTS,
        "n_values": list(range(MAX_N + 1)),
        "settings": {
            "generation_truncations": GENERATION_TRUNCATIONS,
            "generation_working_decimal_digits": GENERATION_WORKING_DIGITS,
            "check_truncation": CHECK_TRUNCATION,
            "check_working_decimal_digits": CHECK_WORKING_DIGITS,
            "bit_guard": BIT_GUARD,
            "bisection_guard_decimal_digits": BISECTION_GUARD_DIGITS,
        },
        "entries": entries,
        "metrics": metrics,
        "independent_checks": controls,
        "limitations": (
            "The finite Hill matrices are solved at high precision and checked "
            "against stronger finite sections and independent identities/shooting "
            "residuals. No rigorous bound for the infinite Fourier tail is proved, "
            "so the computation supports heuristic agreement-checked precision, "
            "not proven enclosures."
        ),
    }
    with open(path, "w", encoding="utf-8") as handle:
        json.dump(payload, handle, indent=2, ensure_ascii=False)
        handle.write("\n")
    return payload


def self_check():
    payload = write_computation(os.devnull)
    print(json.dumps({
        "status": "ok",
        "entries": payload["metrics"]["entries"],
        "minimum_supported_digits_from_stronger_finite_section":
            payload["metrics"]["minimum_supported_digits_from_stronger_finite_section"],
        "worst_check_relative_difference": payload["metrics"]["worst_check_relative_difference"],
    }, indent=2))


def main(argv=None):
    parser = argparse.ArgumentParser()
    parser.add_argument("--write-computation", metavar="PATH",
                        help="write entries and verification metrics as JSON")
    parser.add_argument("--self-check", action="store_true",
                        help="compute every entry and run all offline checks")
    parser.add_argument("--publish", action="store_true",
                        help="publish to NumberDB only together with --allow-publish")
    parser.add_argument("--allow-publish", action="store_true",
                        help="explicit opt-in guard for --publish")
    args = parser.parse_args(argv)

    if args.publish:
        if not args.allow_publish:
            raise SystemExit("--publish requires --allow-publish")
        generator = MathieuCharacteristicValuesA()
        print(generator.publish(message="100-digit Mathieu a_n(q) precision repair"))
        return 0

    if args.write_computation:
        payload = write_computation(args.write_computation)
        print(json.dumps({
            "status": "ok",
            "path": args.write_computation,
            "entries": payload["metrics"]["entries"],
            "minimum_supported_digits_from_stronger_finite_section":
                payload["metrics"]["minimum_supported_digits_from_stronger_finite_section"],
            "worst_check_relative_difference": payload["metrics"]["worst_check_relative_difference"],
        }, indent=2))
        return 0

    if args.self_check:
        self_check()
        return 0

    parser.print_help()
    return 0


if __name__ == "__main__":
    sys.exit(main())