# Copyright (c) 2026 INNOVATIO SAS
# SPDX-License-Identifier: MIT
"""MSC-P-046 (marketing-science-center.com): predict how much a customer will spend per purchase.

Standard library only. The gamma-gamma model of spend per transaction (Fader, Hardie and Lee 2005; Fader and
Hardie 2013) is fitted to the CDNOW sample published by Bruce Hardie: 2,357 customers who made their first
purchase at CDNOW in the first quarter of 1997, followed until the end of June 1998. Weeks 1 to 39 calibrate
the model; weeks 40 to 78 are held out. Before any result, the program checks that it reproduces what the
authors publish on the same data (Table 1, the maximum likelihood estimates, the gap between theoretical and
observed mean, the modes, the correlation between frequency and spend, the log-likelihoods of the 2005
article) and stops otherwise. It then compares three predictions of each customer's future spend per
purchase on the held-out weeks: the customer's own past average, the population mean, and the gamma-gamma
conditional expectation, which weighs the two. The R reference prints exactly the same lines.

The data file CDNOW_sample.txt is read from the current directory, next to the program, or from
editorial/sources/MSC-P-046/ when the program runs inside the repository; it is downloadable from https://www.brucehardie.com/datasets/.
"""
import math
import os

HERE = os.path.dirname(os.path.abspath(__file__))
DATA_CANDIDATES = ["CDNOW_sample.txt", os.path.join(HERE, "CDNOW_sample.txt"), os.path.join(HERE, "..", "..", "sources", "MSC-P-046", "CDNOW_sample.txt")]
RECORDS = 6919
CUSTOMERS = 2357
CALIBRATION_END = 272   # days after 1 January 1997: 30 September 1997, end of week 39
HOLDOUT_END = 545       # 30 June 1998, end of week 78
LCG_SEED = 20261007
LCG_MULTIPLIER = 1103515245
LCG_INCREMENT = 12345
LCG_MODULUS = 2147483648
BOOTSTRAP = 1000
NM_MAX_ITER = 20000
NM_TOL = 1e-9    # the log-likelihood, a sum of some 3,000 logarithms near -4,000, carries rounding noise of about 1e-11
NM_SIZE = 1e-7
SIMPSON_STEPS = 20000
GRADIENT_STEP = 1e-4
GRADIENT_TOL = 1e-3
CV_MIN_PURCHASES = 3
CV_SIMULATIONS = 200
EXAMPLE_SPEND = 100.0   # the authors' customer A: one repeat purchase totalling $100
BANDS = [(0.0, 20.0), (20.0, 50.0), (50.0, 1e9)]

# What the authors publish on these data. Fader and Hardie (2013), Table 1 and p. 6; Fader, Hardie and Lee
# (2005), Table 1, pp. 12, 15 and 17.
PUB_REPEATERS = 946
PUB_TABLE1 = {"minimum": 2.99, "25th percentile": 15.75, "median": 27.50, "75th percentile": 41.80,
              "maximum": 299.63, "mean": 35.08, "standard deviation": 30.28, "mode": 14.96}
PUB_PQG = (6.25, 3.74, 15.44)
PUB_MEAN_GAP_CENTS = 9
PUB_MODEL_MODE = 19
PUB_OBSERVED_MODE = 15
PUB_CORRELATION = 0.11
PUB_CORRELATION_WITHOUT = 0.06
PUB_P_WITHOUT = 0.08
PUB_OUTLIER = (21, 300)
PUB_LL_39 = -4659
PUB_LL_39_AT_78 = -4661
PUB_MEAN_78 = 36
PUB_SKEWNESS = 4
PUB_KURTOSIS = 17


class Declared:
    def __init__(self, seed):
        self.state = seed

    def uniform(self):
        self.state = (LCG_MULTIPLIER * self.state + LCG_INCREMENT) % LCG_MODULUS
        return (self.state + 0.5) / LCG_MODULUS

    def normal(self):
        u1 = self.uniform()
        u2 = self.uniform()
        return math.sqrt(-2.0 * math.log(u1)) * math.cos(2.0 * math.pi * u2)

    def gamma(self, shape):
        """Gamma(shape, 1) for shape >= 1, Marsaglia and Tsang (2000)."""
        d = shape - 1.0 / 3.0
        c = 1.0 / math.sqrt(9.0 * d)
        while True:
            x = self.normal()
            v = 1.0 + c * x
            if v <= 0.0:
                continue
            v = v * v * v
            u = self.uniform()
            if math.log(u) < 0.5 * x * x + d - d * v + d * math.log(v):
                return d * v


def plain_sum(values):
    """Left-to-right double sum; Python 3.12+ sum() compensates, R's sum() uses long double."""
    total = 0.0
    for v in values:
        total += v
    return total


def fmt(x, digits):
    text = f"{x:.{digits}f}"
    return text[1:] if text.startswith("-") and float(text) == 0.0 else text


def pct(x, digits=1):
    return fmt(100.0 * x, digits)


def nearest(x):
    """Nearest integer, halves up (Python's and R's built-in rounding send halves to the even integer)."""
    return math.floor(x + 0.5)


def fail(message):
    raise SystemExit(f"error: {message}")


def lgam(x):
    """log Gamma(x), x > 0, written out identically in Python and R (shift to x >= 10, then Stirling); the libraries' log and exp
    may still differ by one unit in the last place, so the parity holds at the printed precision, not bit for bit."""
    prod = 1.0
    while x < 10.0:
        prod *= x
        x += 1.0
    z = 1.0 / x
    z2 = z * z
    series = z * (1.0 / 12.0 - z2 * (1.0 / 360.0 - z2 * (1.0 / 1260.0 - z2 * (1.0 / 1680.0 - z2 / 1188.0))))
    return (x - 0.5) * math.log(x) - x + 0.5 * math.log(2.0 * math.pi) + series - math.log(prod)


def betacf(a, b, x):
    """Continued fraction of the regularized incomplete beta function (modified Lentz)."""
    tiny = 1e-300
    qab, qap, qam = a + b, a + 1.0, a - 1.0
    c = 1.0
    d = 1.0 - qab * x / qap
    if abs(d) < tiny:
        d = tiny
    d = 1.0 / d
    h = d
    for m in range(1, 301):
        m2 = 2.0 * m
        aa = m * (b - m) * x / ((qam + m2) * (a + m2))
        d = 1.0 + aa * d
        if abs(d) < tiny:
            d = tiny
        c = 1.0 + aa / c
        if abs(c) < tiny:
            c = tiny
        d = 1.0 / d
        h *= d * c
        aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2))
        d = 1.0 + aa * d
        if abs(d) < tiny:
            d = tiny
        c = 1.0 + aa / c
        if abs(c) < tiny:
            c = tiny
        d = 1.0 / d
        delta = d * c
        h *= delta
        if abs(delta - 1.0) < 1e-15:
            break
    return h


def incomplete_beta(a, b, x):
    if x <= 0.0:
        return 0.0
    if x >= 1.0:
        return 1.0
    front = math.exp(lgam(a + b) - lgam(a) - lgam(b) + a * math.log(x) + b * math.log(1.0 - x))
    if x < (a + 1.0) / (a + b + 2.0):
        return front * betacf(a, b, x) / a
    return 1.0 - front * betacf(b, a, 1.0 - x) / b


def t_two_sided(t, df):
    return incomplete_beta(df / 2.0, 0.5, df / (df + t * t))


def nelder_mead(f, start, step):
    k = len(start)
    pts = [list(start)]
    for j in range(k):
        p = list(start)
        p[j] += step
        pts.append(p)
    vals = [f(p) for p in pts]
    converged = False
    for _ in range(NM_MAX_ITER):
        order = sorted(range(k + 1), key=lambda i: (vals[i], i))
        pts = [pts[i] for i in order]
        vals = [vals[i] for i in order]
        size = max(abs(pts[i][j] - pts[0][j]) for i in range(1, k + 1) for j in range(k))
        if vals[k] - vals[0] < NM_TOL and size < NM_SIZE:
            converged = True
            break
        centroid = []
        for j in range(k):
            c = 0.0
            for i in range(k):
                c += pts[i][j]
            centroid.append(c / k)
        refl = [centroid[j] + (centroid[j] - pts[k][j]) for j in range(k)]
        fr = f(refl)
        if fr < vals[0]:
            exp_pt = [centroid[j] + 2.0 * (centroid[j] - pts[k][j]) for j in range(k)]
            fe = f(exp_pt)
            if fe < fr:
                pts[k], vals[k] = exp_pt, fe
            else:
                pts[k], vals[k] = refl, fr
        elif fr < vals[k - 1]:
            pts[k], vals[k] = refl, fr
        else:
            if fr < vals[k]:
                con = [centroid[j] + 0.5 * (refl[j] - centroid[j]) for j in range(k)]
            else:
                con = [centroid[j] + 0.5 * (pts[k][j] - centroid[j]) for j in range(k)]
            fc = f(con)
            if fc < min(fr, vals[k]):
                pts[k], vals[k] = con, fc
            else:
                for i in range(1, k + 1):
                    pts[i] = [pts[0][j] + 0.5 * (pts[i][j] - pts[0][j]) for j in range(k)]
                    vals[i] = f(pts[i])
    best = min(range(k + 1), key=lambda i: (vals[i], i))
    return pts[best], vals[best], converged


# ---- data -------------------------------------------------------------------------------------------------

def day_index(stamp):
    """Days after 1 January 1997 of a YYYYMMDD date (1997 and 1998 only: neither is a leap year)."""
    year, month, day = int(stamp[:4]), int(stamp[4:6]), int(stamp[6:])
    before = [0, 31, 59, 90, 120, 151, 181, 212, 243, 273, 304, 334]
    return 365 * (year - 1997) + before[month - 1] + day - 1


def read_data():
    path = next((p for p in DATA_CANDIDATES if os.path.exists(p)), None)
    if path is None:
        fail("CDNOW_sample.txt not found: download it from https://www.brucehardie.com/datasets/")
    records = []
    with open(path, encoding="ascii") as handle:
        for line in handle:
            fields = line.split()
            if fields:
                records.append((int(fields[1]), day_index(fields[2]), int(fields[3]), float(fields[4])))
    return records


def purchase_days(records):
    """Per customer, the list of (day, dollars) with transactions on the same day added up, in file order."""
    days = {}
    for cid, day, _, dollars in records:
        cust = days.setdefault(cid, [])
        if cust and cust[-1][0] == day:
            cust[-1] = (day, cust[-1][1] + dollars)
        else:
            cust.append((day, dollars))
    return days


def summary_all(days, end):
    """Per customer: number of purchases up to day end, the first one included, and their mean value."""
    out = []
    for cid in range(1, CUSTOMERS + 1):
        values = [v for d, v in days[cid] if d <= end]
        out.append((len(values), plain_sum(values) / len(values)))
    return out


def summary(days, start, end):
    """Per customer: number of repeat purchases with start <= day <= end (the first purchase excluded) and their mean value."""
    out = []
    for cid in range(1, CUSTOMERS + 1):
        values = [v for d, v in days[cid][1:] if start <= d <= end]
        out.append((len(values), plain_sum(values) / len(values) if values else 0.0))
    return out


# ---- the gamma-gamma model --------------------------------------------------------------------------------

class Repeaters:
    """Frequency x and mean spend zbar of the customers with x >= 1, with what the likelihood needs precomputed."""

    def __init__(self, pairs):
        self.x = [x for x, _ in pairs]
        self.z = [z for _, z in pairs]
        self.distinct = sorted(set(self.x))
        self.lx = {k: math.log(k) for k in self.distinct}
        # Customers grouped by x, in their order; per group the count, the sum of log zbar and the zbar values.
        self.groups = []
        for k in self.distinct:
            zs = [z for x, z in zip(self.x, self.z) if x == k]
            self.groups.append((k, float(len(zs)), plain_sum([math.log(z) for z in zs]), zs))


def loglik(theta, data):
    """Log-likelihood of the mean spends, the authors' equation (1a), summed by group of customers with the same x."""
    p, q, g = theta
    lq = lgam(q)
    lng = math.log(g)
    total = 0.0
    for k, count, sum_lz, zs in data.groups:
        logs = 0.0
        for z in zs:
            logs += math.log(g + k * z)
        a = lgam(p * k + q) - lgam(p * k) - lq + q * lng + p * k * data.lx[k]
        total += count * a + (p * k - 1.0) * sum_lz - (p * k + q) * logs
    return total


def fit(data, start, step):
    """Maximum likelihood on log p, log q, log gamma; the simplex is restarted from its best point until it is stable."""
    def objective(v):
        return -loglik([math.exp(v[0]), math.exp(v[1]), math.exp(v[2])], data)
    point, value, _ = nelder_mead(objective, start, step)
    for _ in range(10):
        new_point, new_value, converged = nelder_mead(objective, point, 0.1)
        stable = value - new_value < 1e-9
        point, value = new_point, new_value
        if stable and converged:
            return [math.exp(v) for v in point], -value
    fail("the simplex does not stabilise")


def true_mean_sd(theta):
    """Standard deviation of the customers' true means, p gamma / ((q - 1) sqrt(q - 2)), for q > 2."""
    p, q, g = theta
    return p * g / ((q - 1.0) * math.sqrt(q - 2.0))


def population_mean(theta):
    p, q, g = theta
    return p * g / (q - 1.0)


def own_weight(theta, x):
    """Weight of the customer's own average in the conditional expectation, px / (px + q - 1)."""
    p, q, _ = theta
    return p * x / (p * x + q - 1.0)


def conditional_mean(theta, x, zbar):
    """E(Z | x, zbar), the authors' equation (5); the population mean when x = 0."""
    if x == 0:
        return population_mean(theta)
    p, q, g = theta
    return p * (g + x * zbar) / (p * x + q - 1.0)


def density_zbar(theta, x, zbar):
    """f(zbar | x), form (1b), safe for large x and zbar."""
    p, q, g = theta
    return math.exp(-lgam(p * x) - lgam(q) + lgam(p * x + q) + q * math.log(g / (g + x * zbar)) + p * x * math.log(x * zbar / (g + x * zbar))) / zbar


def model_mode(theta, counts):
    """Mode on the $1 grid 1..300 of the density of zbar averaged over the observed frequencies (the authors' Figure 3)."""
    total = plain_sum([float(c) for c in counts.values()])
    best, best_y = -1.0, 0
    for y in range(1, 301):
        f = plain_sum([counts[k] * density_zbar(theta, k, float(y)) for k in sorted(counts)]) / total
        if f > best:
            best, best_y = f, y
    return best_y


# ---- descriptive statistics -------------------------------------------------------------------------------

def mean_of(rows, key):
    return plain_sum([r[key] for r in rows]) / len(rows)


def mean_abs(rows, key):
    return plain_sum([abs(r["z2"] - r[key]) for r in rows]) / len(rows)


def median(xs):
    s = sorted(xs)
    k = len(s)
    return s[k // 2] if k % 2 else (s[k // 2 - 1] + s[k // 2]) / 2.0


def mean_sd(xs):
    m = plain_sum(xs) / len(xs)
    return m, math.sqrt(plain_sum([(x - m) * (x - m) for x in xs]) / (len(xs) - 1))


def weibull_quantile(sorted_xs, prob):
    """Quantile with plotting position (n + 1) p, R's type 6."""
    h = (len(sorted_xs) + 1) * prob
    lo = math.floor(h)
    return sorted_xs[lo - 1] + (h - lo) * (sorted_xs[lo] - sorted_xs[lo - 1])


def moments(xs):
    m = plain_sum(xs) / len(xs)
    d = [x - m for x in xs]
    m2 = plain_sum([v * v for v in d]) / len(xs)
    m3 = plain_sum([v * v * v for v in d]) / len(xs)
    m4 = plain_sum([(v * v) * (v * v) for v in d]) / len(xs)
    return m3 / (m2 * math.sqrt(m2)), m4 / (m2 * m2) - 3.0


def cents_mode(xs):
    counts = {}
    for x in xs:
        key = math.floor(x * 100.0 + 0.5)
        counts[key] = counts.get(key, 0) + 1
    top = max(counts.values())
    return min(k for k, c in counts.items() if c == top) / 100.0, top


def correlation(xs, ys):
    mx = plain_sum(xs) / len(xs)
    my = plain_sum(ys) / len(ys)
    sxy = plain_sum([(x - mx) * (y - my) for x, y in zip(xs, ys)])
    sxx = plain_sum([(x - mx) * (x - mx) for x in xs])
    syy = plain_sum([(y - my) * (y - my) for y in ys])
    return sxy / math.sqrt(sxx * syy)


def quantile_pair(xs):
    """2.5th and 97.5th sorted values (the 25th and 975th of 1,000)."""
    s = sorted(xs)
    return s[nearest(0.025 * len(s)) - 1], s[nearest(0.975 * len(s)) - 1]


def simpson(f, a, b, steps):
    h = (b - a) / steps
    total = f(a) + f(b)
    for i in range(1, steps):
        total += (4.0 if i % 2 else 2.0) * f(a + i * h)
    return total * h / 3.0


# ---- checks ---------------------------------------------------------------------------------------------------

def check_invariants(theta):
    checks = 0

    def ensure(condition, message):
        nonlocal checks
        checks += 1
        if not condition:
            fail(f"invariant failed: {message}")

    ensure(abs(lgam(1.0)) < 1e-13 and abs(lgam(2.0)) < 1e-13, "log Gamma(1) = log Gamma(2) = 0")
    ensure(abs(lgam(0.5) - 0.5 * math.log(math.pi)) < 1e-13, "log Gamma(1/2) = log(pi) / 2")
    for v in (0.3, 3.7441, 6.2498, 43.75, 131.25):
        ensure(abs(lgam(v) - math.lgamma(v)) < 1e-11 * max(1.0, abs(math.lgamma(v))), f"declared log Gamma({v}) against the library")
    ensure(abs(lgam(7.3) - lgam(6.3) - math.log(6.3)) < 1e-12, "log Gamma(x + 1) = log Gamma(x) + log x")
    ensure(abs(t_two_sided(1.959963984540054, 1e7) - 0.05) < 1e-6, "two-sided p of 1.96 with many degrees of freedom is 0.05")
    ensure(abs(incomplete_beta(2.0, 3.0, 0.4) - 0.5248) < 1e-12, "I_0.4(2, 3) = 0.5248")
    p, q, g = theta
    for x in (1, 4):
        # u = x zbar / (gamma + x zbar) is beta(px, q): the density integrates to 1 and zbar has mean p gamma / (q - 1).
        a, b = p * x, q
        norm = lgam(a + b) - lgam(a) - lgam(b)
        dens = lambda u: 0.0 if u <= 0.0 or u >= 1.0 else math.exp(norm + (a - 1.0) * math.log(u) + (b - 1.0) * math.log(1.0 - u))
        ensure(abs(simpson(dens, 0.0, 1.0, SIMPSON_STEPS) - 1.0) < 1e-8, f"f(zbar | x = {x}) integrates to 1")
        mean = simpson(lambda u: dens(u) * u / (1.0 - u) * g / x if u < 1.0 else 0.0, 0.0, 1.0, SIMPSON_STEPS)
        ensure(abs(mean - population_mean(theta)) < 1e-4, f"E(zbar | x = {x}) equals p gamma / (q - 1)")
        zbar = 27.5
        form_a = math.exp(lgam(a + q) - lgam(a) - lgam(q) + (a - 1.0) * math.log(zbar) + a * math.log(x) + q * math.log(g) - (a + q) * math.log(g + x * zbar))
        ensure(abs(form_a - density_zbar(theta, x, zbar)) < 1e-12 * form_a, f"forms (1a) and (1b) agree at x = {x}")
        w = own_weight(theta, x)
        ensure(abs(conditional_mean(theta, x, zbar) - ((1.0 - w) * population_mean(theta) + w * zbar)) < 1e-10, f"equation (5) is a weighted average at x = {x}")
    ensure(own_weight(theta, 1000000) > 0.9999, "the weight of the own average tends to 1")
    return checks


def check_gradient(theta, data):
    """Stops unless the fitted point is a stationary point: the log-likelihood gradient on log p, log q, log gamma is near zero."""
    point = [math.log(v) for v in theta]
    for j in range(3):
        up = list(point)
        down = list(point)
        up[j] += GRADIENT_STEP
        down[j] -= GRADIENT_STEP
        slope = (loglik([math.exp(v) for v in up], data) - loglik([math.exp(v) for v in down], data)) / (2.0 * GRADIENT_STEP)
        if abs(slope) > GRADIENT_TOL:
            fail("the fitted point is not a maximum of the log-likelihood")


def published_checks(pairs, repeaters, theta, ll, theta78, ll39_at_78, counts):
    z = repeaters.z
    s = sorted(z)
    m, sd = mean_sd(z)
    mode, _ = cents_mode(z)
    table1 = {"minimum": s[0], "25th percentile": weibull_quantile(s, 0.25), "median": weibull_quantile(s, 0.5),
              "75th percentile": weibull_quantile(s, 0.75), "maximum": s[-1], "mean": m, "standard deviation": sd, "mode": mode}
    reproduced = 0
    print("Published values, Fader and Hardie (2013) and Fader, Hardie and Lee (2005), on the same data:")
    print(f"  repeat buyers in weeks 1-39: {len(z)} of {len(pairs)} (published {PUB_REPEATERS} of {CUSTOMERS})")
    if len(z) != PUB_REPEATERS or len(pairs) != CUSTOMERS:
        fail("the number of repeat buyers is not reproduced")
    for name, value in table1.items():
        ok = abs(value - PUB_TABLE1[name]) <= 0.0051
        reproduced += ok
        print(f"  Table 1, {name}: {fmt(value, 4)} (published {fmt(PUB_TABLE1[name], 2)}) {'reproduced' if ok else 'NOT reproduced'}")
    if reproduced != len(table1):
        fail("Table 1 is not reproduced")
    skew, kurt = moments(z)
    print(f"  2005 Table 1, skewness: {fmt(skew, 2)} (published {PUB_SKEWNESS}) {'reproduced' if nearest(skew) == PUB_SKEWNESS else 'NOT reproduced'}")
    print(f"  2005 Table 1, excess kurtosis: {fmt(kurt, 2)} (published {PUB_KURTOSIS}) {'reproduced' if nearest(kurt) == PUB_KURTOSIS else 'NOT reproduced'}")
    if nearest(kurt) != PUB_KURTOSIS:
        fail("the kurtosis is not reproduced")
    for name, value, published in zip(("p", "q", "gamma"), theta, PUB_PQG):
        print(f"  estimate {name}: {fmt(value, 4)} (published {fmt(published, 2)})")
        if abs(value - published) > 0.0051:
            fail(f"the estimate of {name} is not reproduced")
    gap = population_mean(theta) - m
    print(f"  theoretical mean {fmt(population_mean(theta), 4)} minus observed mean {fmt(m, 4)}: {fmt(gap, 4)} dollar (published: {PUB_MEAN_GAP_CENTS} cents)")
    if nearest(100.0 * gap) != PUB_MEAN_GAP_CENTS:
        fail("the gap between theoretical and observed mean is not reproduced")
    model = model_mode(theta, counts)
    print(f"  mode of the model density of zbar: {model} dollars (published {PUB_MODEL_MODE}); observed mode {fmt(mode, 2)} (published {PUB_OBSERVED_MODE})")
    if model != PUB_MODEL_MODE or nearest(mode) != PUB_OBSERVED_MODE:
        fail("the modes are not reproduced")
    xs = [float(x) for x in repeaters.x]
    r = correlation(xs, z)
    outlier = [i for i in range(len(z)) if repeaters.x[i] == PUB_OUTLIER[0]]
    if len(outlier) != 1 or nearest(z[outlier[0]]) != PUB_OUTLIER[1]:
        fail("the outlier (21 transactions, $300) is not found")
    keep = [i for i in range(len(z)) if i != outlier[0]]
    r2 = correlation([xs[i] for i in keep], [z[i] for i in keep])
    df = len(keep) - 2
    pvalue = t_two_sided(r2 * math.sqrt(df) / math.sqrt(1.0 - r2 * r2), float(df))
    print(f"  correlation of frequency and mean spend: {fmt(r, 4)} (published {fmt(PUB_CORRELATION, 2)}); without the customer with {repeaters.x[outlier[0]]} purchases and a mean of {fmt(z[outlier[0]], 2)}: {fmt(r2, 4)}, p = {fmt(pvalue, 4)} (published {fmt(PUB_CORRELATION_WITHOUT, 2)}, p = {fmt(PUB_P_WITHOUT, 2)})")
    if nearest(100.0 * r) != nearest(100.0 * PUB_CORRELATION) or nearest(100.0 * r2) != nearest(100.0 * PUB_CORRELATION_WITHOUT) or nearest(100.0 * pvalue) != nearest(100.0 * PUB_P_WITHOUT):
        fail("the correlations are not reproduced")
    sum_log_x = plain_sum([repeaters.lx[x] for x in repeaters.x])
    print(f"  log-likelihood of the mean spends, weeks 1-39: {fmt(ll, 4)}; minus the sum of log x ({fmt(sum_log_x, 4)}): {fmt(ll - sum_log_x, 4)} (published {PUB_LL_39})")
    print(f"  78-week estimates p = {fmt(theta78[0], 4)}, q = {fmt(theta78[1], 4)}, gamma = {fmt(theta78[2], 4)}; their weeks 1-39 log-likelihood: {fmt(ll39_at_78, 4)}; minus the sum of log x: {fmt(ll39_at_78 - sum_log_x, 4)} (published {PUB_LL_39_AT_78})")
    print(f"  78-week population mean: {fmt(population_mean(theta78), 4)} (published about {PUB_MEAN_78})")
    if nearest(ll - sum_log_x) != PUB_LL_39 or nearest(ll39_at_78 - sum_log_x) != PUB_LL_39_AT_78 or nearest(population_mean(theta78)) != PUB_MEAN_78:
        fail("the 2005 log-likelihoods or the 78-week mean are not reproduced")
    return sum_log_x


# ---- main -------------------------------------------------------------------------------------------------

def main():
    records = read_data()
    customers = sorted({r[0] for r in records})
    if len(records) != RECORDS or customers != list(range(1, CUSTOMERS + 1)):
        fail("CDNOW_sample.txt is not the published file (6,919 records, customers 1 to 2,357)")
    days = purchase_days(records)
    if max(days[c][0][0] for c in customers) > 89:
        fail("a customer's first purchase falls after March 1997")
    purchases = sum(len(v) for v in days.values())
    dollars = plain_sum([r[3] for r in records])
    print("MSC-P-046: gamma-gamma model of spend per purchase, CDNOW sample (Fader and Hardie)")
    print(f"Data: {len(records)} records, {CUSTOMERS} customers, {purchases} purchase days once same-day transactions are added up, {fmt(dollars, 2)} dollars")
    print("Calibration: weeks 1-39 (to 30 September 1997); holdout: weeks 40-78 (1 October 1997 to 30 June 1998)")

    cal = summary(days, 0, CALIBRATION_END)
    full = summary(days, 0, HOLDOUT_END)
    hold = summary(days, CALIBRATION_END + 1, HOLDOUT_END)
    repeaters = Repeaters([c for c in cal if c[0] > 0])
    counts = {}
    for x in repeaters.x:
        counts[x] = counts.get(x, 0) + 1
    theta, ll = fit(repeaters, [0.0, 0.0, 0.0], 1.0)
    full_repeaters = Repeaters([c for c in full if c[0] > 0])
    theta78, _ = fit(full_repeaters, [0.0, 0.0, 0.0], 1.0)
    ll39_at_78 = loglik(theta78, repeaters)

    check_gradient(theta, repeaters)
    published_checks(cal, repeaters, theta, ll, theta78, ll39_at_78, counts)
    print(f"Invariants: {check_invariants(theta)} checks passed")

    p, q, g = theta
    ez = population_mean(theta)
    print("")
    print("Fitted on weeks 1-39 (946 repeat buyers):")
    print(f"  p = {fmt(p, 4)}, q = {fmt(q, 4)}, gamma = {fmt(g, 4)}, log-likelihood {fmt(ll, 2)}")
    print(f"  coefficient of variation of a customer's purchases around their own mean, 1 / sqrt(p): {fmt(1.0 / math.sqrt(p), 4)}")
    print(f"  population mean spend per purchase E(Z) = p gamma / (q - 1): {fmt(ez, 4)} dollars")
    print(f"  spread of the customers' true means, standard deviation: {fmt(true_mean_sd(theta), 4)} dollars")
    print(f"  repeat buyers by number of repeat purchases: " + ", ".join(f"{k}: {counts[k]}" for k in sorted(counts)))

    print("")
    print("Weight of the customer's own average in E(Z | x, zbar), and examples:")
    for x in (1, 2, 3, 4, 6, 8, 10):
        w = own_weight(theta, x)
        w78 = own_weight(theta78, x)
        print(f"  x = {x}: weight {fmt(w, 4)} (78-week fit {fmt(w78, 4)}); zbar 20 -> {fmt(conditional_mean(theta, x, 20.0), 2)}, zbar 50 -> {fmt(conditional_mean(theta, x, 50.0), 2)}, zbar 100 -> {fmt(conditional_mean(theta, x, 100.0), 2)}")
    first90 = next(x for x in range(1, 100) if own_weight(theta, x) >= 0.9)
    first90_78 = next(x for x in range(1, 100) if own_weight(theta78, x) >= 0.9)
    print(f"  repeat purchases before the own average weighs 90 %: {first90} (78-week fit: {first90_78})")
    print(f"  customer A, one repeat purchase of {fmt(EXAMPLE_SPEND, 0)} dollars: E(Z) = {fmt(conditional_mean(theta, 1, EXAMPLE_SPEND), 2)} dollars")

    print(f"  customer A with the 78-week fit: E(Z) = {fmt(conditional_mean(theta78, 1, EXAMPLE_SPEND), 2)} dollars")

    # The page's variant: the same model fitted on all purchases of weeks 1-39, the first one included.
    alls = summary_all(days, CALIBRATION_END)
    all_fit_set = Repeaters([c for c in alls if c[1] > 0.0])
    theta_all, ll_all = fit(all_fit_set, [0.0, 0.0, 0.0], 1.0)
    print("")
    print("Variant of the page: gamma-gamma fitted on all purchases of weeks 1-39, the first one included:")
    print(f"  customers with a mean of zero left out of the fit: {CUSTOMERS - len(all_fit_set.x)}")
    print(f"  p = {fmt(theta_all[0], 4)}, q = {fmt(theta_all[1], 4)}, gamma = {fmt(theta_all[2], 4)}, log-likelihood {fmt(ll_all, 2)}; population mean {fmt(population_mean(theta_all), 4)} dollars")
    first_values = [days[c][0][1] for c in range(1, CUSTOMERS + 1)]
    repeat_values = [v for c in range(1, CUSTOMERS + 1) for d, v in days[c][1:] if d <= CALIBRATION_END]
    holdout_values = [v for c in range(1, CUSTOMERS + 1) for d, v in days[c] if CALIBRATION_END < d <= HOLDOUT_END]
    print(f"  mean value of a first purchase {fmt(plain_sum(first_values) / len(first_values), 2)}, of a repeat purchase in weeks 1-39 {fmt(plain_sum(repeat_values) / len(repeat_values), 2)}, of a purchase in weeks 40-78 {fmt(plain_sum(holdout_values) / len(holdout_values), 2)} ({len(holdout_values)} purchases)")
    first_rep = [days[c][0][1] for c in range(1, CUSTOMERS + 1) if cal[c - 1][0] > 0]
    first_none = [days[c][0][1] for c in range(1, CUSTOMERS + 1) if cal[c - 1][0] == 0]
    ratios = [days[c][0][1] / cal[c - 1][1] for c in range(1, CUSTOMERS + 1) if cal[c - 1][0] > 0]
    print(f"  first purchase of the {len(first_rep)} repeat buyers {fmt(plain_sum(first_rep) / len(first_rep), 2)} against the mean of their repeat purchases {fmt(plain_sum(repeaters.z) / len(repeaters.z), 2)} (median ratio {fmt(median(ratios), 4)}); first purchase of the {len(first_none)} customers without repeat purchase {fmt(plain_sum(first_none) / len(first_none), 2)}")
    m_rep, sd_rep = mean_sd(first_rep)
    m_none, sd_none = mean_sd(first_none)
    se_gap = math.sqrt(sd_rep * sd_rep / len(first_rep) + sd_none * sd_none / len(first_none))
    print(f"  first purchase, repeat buyers minus customers without repeat purchase: {fmt(m_rep - m_none, 2)} (standard error {fmt(se_gap, 2)})")
    print(f"  one purchase of {fmt(EXAMPLE_SPEND, 0)} dollars: {fmt(conditional_mean(theta_all, 1, EXAMPLE_SPEND), 2)}; two purchases averaging {fmt(EXAMPLE_SPEND, 0)}: {fmt(conditional_mean(theta_all, 2, EXAMPLE_SPEND), 2)}; weight of one purchase {fmt(own_weight(theta_all, 1), 4)}, of two {fmt(own_weight(theta_all, 2), 4)}")

    # Holdout: what each customer actually spent per purchase in weeks 40-78.
    print("")
    print("Holdout, weeks 40-78 (error = actual mean spend per purchase minus prediction; standard errors across customers):")
    rows = []
    for (x, zbar), (n_all, z_all), (n2, z2) in zip(cal, alls, hold):
        rows.append({"x": x, "zbar": zbar, "n_all": n_all, "z_all": z_all, "n2": n2, "z2": z2,
                     "rep": zbar, "all": z_all, "pop": ez, "gg": conditional_mean(theta, x, zbar),
                     "gg_all": conditional_mean(theta_all, n_all, z_all)})
    labels = {"rep": "average of repeat purchases", "all": "average of all purchases", "pop": "population mean",
              "gg": "gamma-gamma, repeat purchases (authors)", "gg_all": "gamma-gamma, all purchases (page)"}
    sets = [("repeat buyers of weeks 1-39 who buy again", [r for r in rows if r["x"] > 0 and r["n2"] > 0], ["rep", "all", "pop", "gg", "gg_all"]),
            ("customers without repeat purchase in weeks 1-39 who buy", [r for r in rows if r["x"] == 0 and r["n2"] > 0], ["all", "pop", "gg_all"]),
            ("all customers who buy in weeks 40-78", [r for r in rows if r["n2"] > 0], ["all", "gg", "gg_all"])]
    for label, test, keys in sets:
        print(f"  {label}: {len(test)}")
        for key in keys:
            e = [r["z2"] - r[key] for r in test]
            mae = plain_sum([abs(v) for v in e]) / len(e)
            rmse = math.sqrt(plain_sum([v * v for v in e]) / len(e))
            bias, bias_sd = mean_sd(e)
            weighted = plain_sum([r["n2"] * abs(r["z2"] - r[key]) for r in test]) / plain_sum([float(r["n2"]) for r in test])
            line = f"    {labels[key]}: mean absolute error {fmt(mae, 2)}, root mean square error {fmt(rmse, 2)}, mean error {fmt(bias, 2)} (standard error {fmt(bias_sd / math.sqrt(len(e)), 2)}), mean absolute error weighted by purchases {fmt(weighted, 2)}"
            if key != "all":
                d = [abs(r["z2"] - r[key]) - abs(r["z2"] - r["all"]) for r in test]
                dm, dsd = mean_sd(d)
                line += f"; minus average of all purchases {fmt(dm, 2)} (standard error {fmt(dsd / math.sqrt(len(d)), 2)})"
            if key not in ("gg", "pop"):
                d = [abs(r["z2"] - r[key]) - abs(r["z2"] - r["gg"]) for r in test]
                dm, dsd = mean_sd(d)
                line += f"; minus authors' model {fmt(dm, 2)} (standard error {fmt(dsd / math.sqrt(len(d)), 2)})"
            print(line)
    test = sets[0][1]
    better = sum(1 for r in test if abs(r["z2"] - r["gg_all"]) < abs(r["z2"] - r["all"]))
    better_rep = sum(1 for r in test if abs(r["z2"] - r["gg"]) < abs(r["z2"] - r["rep"]))
    print(f"  repeat buyers for whom the page's variant beats the average of all purchases: {better} of {len(test)} ({pct(better / len(test))} %); the authors' model beats the average of repeat purchases for {better_rep} ({pct(better_rep / len(test))} %)")
    print(f"  authors' model against average of repeat purchases: mean absolute error {pct(mean_abs(test, 'gg') / mean_abs(test, 'rep') - 1.0)} %; page's variant against average of all purchases: {pct(mean_abs(test, 'gg_all') / mean_abs(test, 'all') - 1.0)} %")
    print("  by average of repeat purchases in weeks 1-39 (customers, mean zbar, mean prediction of the authors' model, of the page's variant, actual mean spend in weeks 40-78, actual minus authors' prediction with its standard error):")
    for lo, hi in BANDS:
        band = [r for r in test if lo <= r["zbar"] < hi]
        label = f"{fmt(lo, 0)} to {fmt(hi, 0)}" if hi < 1e8 else f"{fmt(lo, 0)} and more"
        dm, dsd = mean_sd([r["z2"] - r["gg"] for r in band])
        print(f"    {label}: {len(band)}, {fmt(mean_of(band, 'zbar'), 2)}, {fmt(mean_of(band, 'gg'), 2)}, {fmt(mean_of(band, 'gg_all'), 2)}, {fmt(mean_of(band, 'z2'), 2)}, {fmt(dm, 2)} ({fmt(dsd / math.sqrt(len(band)), 2)})")
    print("  by repeat purchases in weeks 1-39, mean total spend in weeks 40-78 per customer (customers, actual, authors' model times actual purchases, difference with its standard error):")
    groups = [(k, k) for k in range(7)] + [(7, 1000000)]
    for lo, hi in groups:
        grp = [r for r in rows if lo <= r["x"] <= hi]
        actual = plain_sum([r["n2"] * r["z2"] for r in grp]) / len(grp)
        expected = plain_sum([r["n2"] * r["gg"] for r in grp]) / len(grp)
        dm, dsd = mean_sd([r["n2"] * (r["z2"] - r["gg"]) for r in grp])
        print(f"    {lo if lo == hi else str(lo) + '+'}: {len(grp)}, {fmt(actual, 2)}, {fmt(expected, 2)}, {fmt(dm, 2)} ({fmt(dsd / math.sqrt(len(grp)), 2)})")
    actual_total = plain_sum([r["n2"] * r["z2"] for r in rows])
    for key in ("gg", "gg_all", "all"):
        total = plain_sum([r["n2"] * r[key] for r in rows])
        print(f"  total holdout spend, {labels[key]} times actual purchases: {fmt(total, 2)} against {fmt(actual_total, 2)} ({pct(total / actual_total - 1.0)} %)")

    # Does a customer's spend vary around their mean as much as the model says (coefficient of variation 1 / sqrt(p))?
    stream_cv = Declared(LCG_SEED + 1)
    cvs = []
    for c in range(1, CUSTOMERS + 1):
        values = [v for d, v in days[c][1:] if d <= CALIBRATION_END]
        if len(values) >= CV_MIN_PURCHASES:
            m, sd = mean_sd(values)
            cvs.append((len(values), m, sd / m))
    observed = median([cv for _, _, cv in cvs])
    medians = []
    for _ in range(CV_SIMULATIONS):
        sim = []
        for k, _, _ in cvs:
            draw = [stream_cv.gamma(p) for _ in range(k)]
            m, sd = mean_sd(draw)
            sim.append(sd / m)
        medians.append(median(sim))
    lo, hi = quantile_pair(medians)
    print("")
    print(f"Spread of a customer's purchases around their own mean, repeat buyers with at least {CV_MIN_PURCHASES} repeat purchases in weeks 1-39 ({len(cvs)} customers):")
    print(f"  median observed coefficient of variation {fmt(observed, 4)}; expected under the model with the same numbers of purchases, mean of {CV_SIMULATIONS} simulations {fmt(plain_sum(medians) / len(medians), 4)}, 95 % of simulations between {fmt(lo, 4)} and {fmt(hi, 4)}")
    for lo_b, hi_b in BANDS:
        band = [cv for _, m, cv in cvs if lo_b <= m < hi_b]
        label = f"{fmt(lo_b, 0)} to {fmt(hi_b, 0)}" if hi_b < 1e8 else f"{fmt(lo_b, 0)} and more"
        print(f"    mean spend {label}: {len(band)} customers, median coefficient of variation {fmt(median(band), 4)}")

    # Uncertainty: customers resampled with replacement, model refitted each time.
    stream = Declared(LCG_SEED)
    n = len(repeaters.x)
    draws = {"p": [], "q": [], "gamma": [], "E(Z)": [], "customer A": [], "weight x = 1": [], "coefficient of variation 1 / sqrt(p)": [], "standard deviation of the true means": []}
    pairs = list(zip(repeaters.x, repeaters.z))
    start = [math.log(v) for v in theta]
    for _ in range(BOOTSTRAP):
        sample = Repeaters([pairs[int(stream.uniform() * n)] for _ in range(n)])
        th, _ = fit(sample, start, 0.1)
        draws["p"].append(th[0])
        draws["q"].append(th[1])
        draws["gamma"].append(th[2])
        draws["E(Z)"].append(population_mean(th))
        draws["customer A"].append(conditional_mean(th, 1, EXAMPLE_SPEND))
        draws["weight x = 1"].append(own_weight(th, 1))
        draws["coefficient of variation 1 / sqrt(p)"].append(1.0 / math.sqrt(th[0]))
        draws["standard deviation of the true means"].append(true_mean_sd(th) if th[1] > 2.0 else 1e9)
    print("")
    print(f"Bootstrap, {BOOTSTRAP} resamples of the {n} repeat buyers, authors' model refitted (seed {LCG_SEED}), nominal 95 % interval:")
    for name, values in draws.items():
        lo, hi = quantile_pair(values)
        print(f"  {name}: {fmt(lo, 4)} to {fmt(hi, 4)}")

if __name__ == "__main__":
    main()
