"""Independent reproduction from the study's mathematical specifications.

The reference is an exact integer sum of the dyadic input numerators, converted
once to binary64 with ties-to-even rounding and scaled by an exact power of two.
No supplied implementation or declared result is read by this program.
"""
import json
import math
import pathlib
import random
import statistics
import sys

LENGTHS = (1000, 10000, 100000, 1000000)
TRIALS = 50
DENOMINATOR = 1 << 53
SCALE = 2.0 ** -53


def sequential(xs):
    accumulator = 0.0
    for x in xs:
        accumulator = accumulator + x
    return accumulator


def adjacent_tree(xs):
    buffer = xs.copy()
    count = len(buffer)
    while count > 1:
        output = 0
        index = 0
        while index + 1 < count:
            buffer[output] = buffer[index] + buffer[index + 1]
            output += 1
            index += 2
        if index < count:
            buffer[output] = buffer[index]
            output += 1
        count = output
    return buffer[0] if count else 0.0


def compensated(xs):
    accumulator = 0.0
    correction = 0.0
    for x in xs:
        adjusted = x - correction
        updated = accumulator + adjusted
        correction = (updated - accumulator) - adjusted
        accumulator = updated
    return accumulator


def magnitude_compensated(xs):
    accumulator = 0.0
    correction = 0.0
    for x in xs:
        updated = accumulator + x
        if abs(accumulator) >= abs(x):
            correction += (accumulator - updated) + x
        else:
            correction += (x - updated) + accumulator
        accumulator = updated
    return accumulator + correction


ALGORITHMS = {
    "recursive": sequential,
    "pairwise": adjacent_tree,
    "kahan": compensated,
    "neumaier": magnitude_compensated,
}


def main():
    aggregate = {"lengths": list(LENGTHS), "trials": TRIALS}
    audit = {"reference": "exact dyadic integer numerators, one ties-to-even conversion", "fsum_agreements": 0, "standard_medians": {}}
    for distribution_code, distribution in enumerate(("positive", "mixed")):
        by_n = {}
        audit["standard_medians"][distribution] = {}
        for n in LENGTHS:
            ulps = {method: [] for method in ALGORITHMS}
            scaled_ulps = {method: [] for method in ALGORITHMS}
            conditions = []
            for trial in range(TRIALS):
                rng = random.Random(2 * (TRIALS * n + trial) + distribution_code)
                numerators = [int(rng.random() * DENOMINATOR) for _ in range(n)]
                if distribution_code:
                    numerators = [2 * numerator - DENOMINATOR for numerator in numerators]
                values = [numerator * SCALE for numerator in numerators]
                exact = float(sum(numerators)) * SCALE
                magnitude = float(sum(abs(x) for x in numerators)) * SCALE
                if math.fsum(values) != exact or math.fsum(abs(x) for x in values) != magnitude:
                    raise AssertionError("math.fsum differs from the independently rounded exact dyadic reference")
                audit["fsum_agreements"] += 1
                conditions.append(magnitude / abs(exact))
                for method, algorithm in ALGORITHMS.items():
                    error = abs(algorithm(values) - exact)
                    ulps[method].append(error / math.ulp(exact))
                    scaled_ulps[method].append(error / math.ulp(magnitude))
            by_n[str(n)] = {
                method: {
                    "mean_ulps": round(statistics.fmean(ulps[method]), 3),
                    "max_ulps": max(ulps[method]),
                    "exact": ulps[method].count(0.0),
                    "mean_ulps_of_magnitudes": round(statistics.fmean(scaled_ulps[method]), 6),
                }
                for method in ALGORITHMS
            }
            # Reproduce the study's stated implementation convention, while
            # separately recording the conventional median of an even sample.
            by_n[str(n)]["median_condition_number"] = round(sorted(conditions)[TRIALS // 2], 3)
            audit["standard_medians"][distribution][str(n)] = statistics.median(conditions)
            print(f"Completed {distribution}, n={n}, trials={TRIALS}", file=sys.stderr, flush=True)
        slopes = {}
        for method in ("recursive", "pairwise"):
            means = [by_n[str(n)][method]["mean_ulps"] for n in LENGTHS]
            if min(means) > 0:
                slopes[method] = round(statistics.linear_regression(
                    [math.log10(n) for n in LENGTHS],
                    [math.log10(mean) for mean in means],
                ).slope, 3)
        aggregate[distribution] = {
            "by_n": by_n,
            "growth_exponent": slopes,
            "max_ulps": {method: max(by_n[str(n)][method]["max_ulps"] for n in LENGTHS) for method in ALGORITHMS},
        }
    aggregate["trials_total"] = 2 * len(LENGTHS) * TRIALS
    aggregate["exact_trials"] = {
        method: sum(aggregate[d]["by_n"][str(n)][method]["exact"] for d in ("positive", "mixed") for n in LENGTHS)
        for method in ALGORITHMS
    }
    pathlib.Path("results/independent.json").write_text(json.dumps(aggregate, indent=2) + "\n")
    pathlib.Path("results/reference-audit.json").write_text(json.dumps(audit, indent=2) + "\n")


if __name__ == "__main__":
    main()
