"""refit_ebbinghaus_c6733.py — re-fit Ebbinghaus's 1885 savings table with three curve shapes.

DATA: Ebbinghaus, "Memory" (1885; Ruger & Bussenius translation 1913), Chapter VII, Section 29, summary table
(psychclassics.yorku.ca/Ebbinghaus/memory7.htm, image table23.jpg, archived in sources/ebb_table23.jpg):
  interval (hours): 0.33, 1, 8.8, 24, 48, 6x24, 31x24
  savings Q (%):    58.2, 44.2, 35.8, 33.7, 27.8, 25.4, 21.1
His own formula (same section, sources/ebb_table24.jpg): b = 100k / ((log10 t)^c + k), t in minutes, k = 1.84, c = 1.25,
fitted by him "with merely approximate estimates, not involving exact calculation by the method of least squares".

Fits (least squares on the seven points, t in minutes):
  log    b = 100k / ((log10 t)^c + k)      — his form, re-fitted and also evaluated at his own k, c
  power  b = a * t^(-d)                     — the form most modern re-analyses prefer (Wixted & Ebbesen 1991)
  expo   b = a * exp(-t / s)                — the smooth curve the training-industry diagrams draw
Prints each fit's RMSE and what it implies at 7 days and 30 days. No network. Run: python refit_ebbinghaus_c6733.py
"""
import math

# the EXACT intervals in minutes from his formula table (ebb_table24.jpg, column t): the summary table's "1 hour" was
# really 64 minutes and "8.8 hours" 526 minutes. Using the rounded hours misses his own Calculated column at t=64.
T_MIN = [20, 64, 526, 1440, 2 * 1440, 6 * 1440, 31 * 1440]
Q = [58.2, 44.2, 35.8, 33.7, 27.8, 25.4, 21.1]


def rmse(f):
    return math.sqrt(sum((f(t) - q) ** 2 for t, q in zip(T_MIN, Q)) / len(Q))


def grid_fit(make, grids):
    """Brute-force least squares over parameter grids (deterministic, no scipy), refined twice around the best point."""
    best = None
    for _ in range(3):
        for p in _product(grids):
            e = rmse(make(*p))
            if best is None or e < best[0]:
                best = (e, p)
        grids = [_refine(g, v) for g, v in zip(grids, best[1])]
    return best


def _product(grids):
    if not grids:
        yield ()
        return
    for v in grids[0]:
        for rest in _product(grids[1:]):
            yield (v,) + rest


def _refine(g, v):
    step = (g[-1] - g[0]) / (len(g) - 1) if len(g) > 1 else 1
    lo, hi = v - 2 * step, v + 2 * step
    n = len(g)
    return [lo + (hi - lo) * i / (n - 1) for i in range(n)]


def lin(lo, hi, n):
    return [lo + (hi - lo) * i / (n - 1) for i in range(n)]


def log_form(k, c):
    return lambda t: 100 * k / (math.log10(t) ** c + k)


def power(a, d):
    return lambda t: a * t ** (-d)


def expo(a, s):
    return lambda t: a * math.exp(-t / s)


def main():
    week, month = 7 * 24 * 60, 30 * 24 * 60
    rows = []
    his = log_form(1.84, 1.25)
    rows.append(("log, his k=1.84 c=1.25", rmse(his), his, "k=1.84 c=1.25"))
    e, (k, c) = grid_fit(log_form, [lin(0.5, 4, 36), lin(0.5, 2.5, 41)])
    rows.append(("log, re-fitted", e, log_form(k, c), "k=%.3f c=%.3f" % (k, c)))
    e, (a, d) = grid_fit(power, [lin(40, 140, 51), lin(0.01, 0.4, 40)])
    rows.append(("power a*t^-d", e, power(a, d), "a=%.2f d=%.4f" % (a, d)))
    e, (a, s) = grid_fit(expo, [lin(20, 80, 61), lin(100, 60000, 120)])
    rows.append(("exponential a*e^(-t/s)", e, expo(a, s), "a=%.2f s=%.0f min" % (a, s)))
    print("observed savings %:", dict(zip(["20m", "1h", "8.8h", "1d", "2d", "6d", "31d"], Q)))
    print("%-26s %8s %9s %9s  %s" % ("model", "RMSE", "b(7 d)", "b(30 d)", "params"))
    for name, e, f, p in rows:
        print("%-26s %8.2f %9.1f %9.1f  %s" % (name, e, f(week), f(month), p))
    # integrity check: his own "Calculated" column (sources/ebb_table24.jpg) must come out of his formula as coded here
    his_calc = [57.0, 46.7, 34.5, 30.4, 28.1, 24.9, 21.2]
    ours = [round(his(t), 1) for t in T_MIN]
    worst = max(abs(a - b) for a, b in zip(ours, his_calc))
    print("\nhis Calculated column reproduced to within %.1f point(s): ours %s | his %s  (t=64 gives %.2f; he printed 46.7 —"
          " a hand calculation from log tables)" % (worst, ours, his_calc, his(64)))
    print("\n'forgotten' in Ebbinghaus's own column IV is 100 - Q: at 1 d %.1f, 6 d %.1f, 31 d %.1f (never 90)"
          % (100 - Q[3], 100 - Q[5], 100 - Q[6]))


if __name__ == "__main__":
    main()
