"""How far four ways of adding up doubles land from the correctly rounded sum.

For each length n in LENGTHS and each of TRIALS seeds, draws n doubles from Python's Mersenne
Twister, uniform on [0, 1) and, separately, on [-1, 1), and sums them four ways: left to right
(recursive summation), pairwise, with Kahan's compensation, and with Neumaier's. Each result is
compared with math.fsum, which returns the correctly rounded sum, and the error is counted in
units in the last place (ulps) of that sum, and in ulps of the sum of the values' magnitudes,
the scale error bounds for summation are stated in, which stays meaningful when the sum nearly
cancels. Writes results/R1.json.

random.Random(seed).random() is the one generator Python promises to reproduce across versions,
and floating-point addition and subtraction are rounded the same way on every IEEE 754 machine,
so every error here is the same wherever the code runs. Python's built-in sum isn't used: it
compensates since Python 3.12, so it would sum differently by version.
"""

import json
import math
import pathlib
import random

LENGTHS = (10**3, 10**4, 10**5, 10**6)
TRIALS = 50  # trials 0, 1, ..., TRIALS - 1 for each length and distribution


def recursive(values):
    total = 0.0
    for value in values:
        total += value
    return total


def pairwise(values):
    """Adds neighbors, then neighbors of those sums, and so on, as a balanced binary tree."""
    level = list(values)
    while len(level) > 1:
        paired = [a + b for a, b in zip(level[0::2], level[1::2])]
        if len(level) % 2:
            paired.append(level[-1])
        level = paired
    return level[0] if level else 0.0


def kahan(values):
    total = 0.0
    carry = 0.0  # the low-order part lost from total so far, negated
    for value in values:
        y = value - carry
        t = total + y
        carry = (t - total) - y
        total = t
    return total


def neumaier(values):
    total = 0.0
    lost = 0.0  # the low-order parts lost so far, added back at the end
    for value in values:
        t = total + value
        if abs(total) >= abs(value):
            lost += (total - t) + value
        else:
            lost += (value - t) + total
        total = t
    return total + lost


METHODS = {"recursive": recursive, "pairwise": pairwise, "kahan": kahan, "neumaier": neumaier}
DISTRIBUTIONS = {"positive": lambda r: r.random(), "mixed": lambda r: 2.0 * r.random() - 1.0}


def draw(distribution, n, seed):
    """n values from the generator seeded with an integer that names the distribution, n, and seed."""
    rng = random.Random(2 * (n * TRIALS + seed) + list(DISTRIBUTIONS).index(distribution))
    sample = DISTRIBUTIONS[distribution]
    return [sample(rng) for _ in range(n)]


def slope(xs, ys):
    """The least-squares slope of log10(y) against log10(x)."""
    lx = [math.log10(x) for x in xs]
    ly = [math.log10(y) for y in ys]
    mx, my = math.fsum(lx) / len(lx), math.fsum(ly) / len(ly)
    return math.fsum((a - mx) * (b - my) for a, b in zip(lx, ly)) / math.fsum((a - mx) ** 2 for a in lx)


def main():
    results = {"lengths": list(LENGTHS), "trials": TRIALS}
    for distribution in DISTRIBUTIONS:
        by_n = {}
        for n in LENGTHS:
            errors = {method: [] for method in METHODS}
            scaled = {method: [] for method in METHODS}
            conditioning = []
            for seed in range(TRIALS):
                values = draw(distribution, n, seed)
                exact = math.fsum(values)
                magnitude = math.fsum(abs(v) for v in values)
                # How ill-conditioned the sum is: the sum of magnitudes over the magnitude of the sum.
                conditioning.append(magnitude / abs(exact))
                for method, add in METHODS.items():
                    error = abs(add(values) - exact)
                    errors[method].append(error / math.ulp(exact))
                    scaled[method].append(error / math.ulp(magnitude))
            by_n[str(n)] = {
                method: {
                    "mean_ulps": round(math.fsum(e) / TRIALS, 3),
                    "max_ulps": max(e),
                    "exact": sum(1 for x in e if x == 0),
                    "mean_ulps_of_magnitudes": round(math.fsum(scaled[method]) / TRIALS, 6),
                }
                for method, e in errors.items()
            }
            by_n[str(n)]["median_condition_number"] = round(sorted(conditioning)[TRIALS // 2], 3)
        growth = {
            method: round(slope(LENGTHS, [by_n[str(n)][method]["mean_ulps"] for n in LENGTHS]), 3)
            for method in ("recursive", "pairwise")
            if all(by_n[str(n)][method]["mean_ulps"] > 0 for n in LENGTHS)
        }
        # The largest error each method made at any length, in ulps of the correctly rounded sum.
        largest = {method: max(by_n[str(n)][method]["max_ulps"] for n in LENGTHS) for method in METHODS}
        results[distribution] = {"by_n": by_n, "growth_exponent": growth, "max_ulps": largest}
    # Over both distributions and every length: how many trials each method summed exactly.
    results["trials_total"] = 2 * len(LENGTHS) * TRIALS
    results["exact_trials"] = {
        method: sum(results[d]["by_n"][str(n)][method]["exact"] for d in DISTRIBUTIONS for n in LENGTHS) for method in METHODS
    }
    out = pathlib.Path(__file__).resolve().parent.parent / "results" / "R1.json"
    out.write_text(json.dumps(results, indent=2) + "\n")
    print(json.dumps(results, indent=2))


if __name__ == "__main__":
    main()
