# Copyright (c) 2026 INNOVATIO SAS
# SPDX-License-Identifier: MIT
"""MSC-P-044 (marketing-science-center.com): value customers who can only buy at fixed occasions, with BG/BB.

Standard library only. Discrete-time noncontractual setting: at each transaction opportunity (one
ski season) a customer buys a season pass or not, and the firm never sees a customer leave. The
beta-geometric/beta-Bernoulli (BG/BB) model of Fader, Hardie and Shang (2010) is checked first on the
published donation data (their Table 2), then applied to a synthetic cohort, declared as such: 4,000
first-time season-pass buyers whose purchase probability p follows Beta(TRUE_ALPHA, TRUE_BETA) and whose
dropout probability theta follows Beta(TRUE_GAMMA, TRUE_DELTA). Six seasons are used for calibration and
five more are kept as a holdout. Values use a margin of 150 per pass and an annual discount rate of 10 %,
the first future pass being discounted by one season. The shortcuts compared with BG/BB are written out in
full: the extended rhythm and the inactivity rule value the past rhythm for ever, with no departure; the damped
rhythm lets it decay at the cohort's observed buyer retention. The R reference prints exactly the same lines.
"""
import math

LCG_SEED = 20261005
LCG_MULTIPLIER = 1103515245
LCG_INCREMENT = 12345
LCG_MODULUS = 2147483648

CUSTOMERS = 4000
TRUE_ALPHA = 1.0
TRUE_BETA = 0.8
TRUE_GAMMA = 0.5
TRUE_DELTA = 2.5
CALIBRATION = 6
HOLDOUT = 5
MARGIN = 150.0
DISCOUNT = 0.10
TERMS = 2000
BOOTSTRAP = 1000
REPLICATIONS = 100
NM_MAX_ITER = 20000
NM_TOL = 1e-10
# Far from any optimum (a parameter above e^50), the log-likelihood is not evaluated: the simplex is told
# "much worse" instead of meeting an overflow. The same bound is used in R.
LOG_PARAMETER_MAX = 50.0
PENALTY = 1e300

# Fader, Hardie and Shang (2010), Table 2: 1995 cohort, n = 6; (x, t_x, number of donors).
DONORS = [(6, 6, 1203), (5, 6, 728), (4, 6, 512), (3, 6, 357), (2, 6, 234), (1, 6, 129),
          (5, 5, 335), (4, 5, 284), (3, 5, 225), (2, 5, 173), (1, 5, 119),
          (4, 4, 240), (3, 4, 181), (2, 4, 155), (1, 4, 78),
          (3, 3, 322), (2, 3, 255), (1, 3, 129),
          (2, 2, 613), (1, 2, 277), (1, 1, 1091), (0, 0, 3464)]


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 patterns(n):
    """The n(n + 1)/2 + 1 recency/frequency patterns (x, t_x), in the order of Table 2."""
    out = []
    for tx in range(n, 0, -1):
        for x in range(tx, 0, -1):
            out.append((x, tx))
    out.append((0, 0))
    return out


def ratio_ab(a, b, x, m):
    """B(a + x, b + m - x) / B(a, b): x purchases in m opportunities, by ascending products."""
    num = 1.0
    for k in range(x):
        num *= a + k
    for k in range(m - x):
        num *= b + k
    den = 1.0
    for k in range(m):
        den *= a + b + k
    return num / den


def alive_through(g, d, t):
    """B(g, d + t) / B(g, d): still alive after t opportunities."""
    s = 1.0
    for k in range(t):
        s *= (d + k) / (g + d + k)
    return s


def dies_at(g, d, t):
    """B(g + 1, d + t) / B(g, d): alive through t opportunities, dead at the start of opportunity t + 1."""
    return alive_through(g, d, t) * g / (g + d + t)


def likelihood(par, x, tx, n):
    """Equation (5) of the paper, written with ascending products instead of beta functions."""
    a, b, g, d = par
    total = ratio_ab(a, b, x, n) * alive_through(g, d, n)
    for i in range(n - tx):
        total += ratio_ab(a, b, x, tx + i) * dies_at(g, d, tx + i)
    return total


def loglik(par, data, n):
    total = 0.0
    for x, tx, f in data:
        if f > 0:
            total += f * math.log(likelihood(par, x, tx, n))
    return total


def bb_loglik(a, b, data, n):
    """Beta-Bernoulli: no dropout, every customer alive for ever."""
    total = 0.0
    for x, _, f in data:
        if f > 0:
            total += f * math.log(ratio_ab(a, b, x, n))
    return total


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 < 1e-8:
            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


def fit_bgbb(data, n, start=(0.0, 0.0, 0.0, 0.0)):
    """Maximum likelihood on the log scale, restarted from its own optimum until it stops moving."""
    def neg(z):
        if max(z) > LOG_PARAMETER_MAX:
            return PENALTY
        return -loglik([math.exp(v) for v in z], data, n)

    z = list(start)
    value = neg(z)
    for _ in range(20):
        z_new, v_new, converged = nelder_mead(neg, z, 0.5)
        if not converged:
            raise SystemExit("error: the BG/BB fit did not converge within NM_MAX_ITER iterations")
        moved = value - v_new
        z, value = z_new, v_new
        if moved < 1e-9:
            return [math.exp(v) for v in z], -value
    raise SystemExit("error: the BG/BB fit kept moving after 20 restarts")


def fit_bb(data, n):
    def neg(z):
        if max(z) > LOG_PARAMETER_MAX:
            return PENALTY
        return -bb_loglik(math.exp(z[0]), math.exp(z[1]), data, n)

    z = [0.0, 0.0]
    value = neg(z)
    for _ in range(20):
        z_new, v_new, converged = nelder_mead(neg, z, 0.5)
        if not converged:
            raise SystemExit("error: the BB fit did not converge within NM_MAX_ITER iterations")
        moved = value - v_new
        z, value = z_new, v_new
        if moved < 1e-9:
            return math.exp(z[0]), math.exp(z[1]), -value
    raise SystemExit("error: the BB fit kept moving after 20 restarts")


def expected_next(par, x, tx, n, horizon):
    """Equation (13): expected purchases over the next `horizon` opportunities, as a sum of per-opportunity terms."""
    a, b, g, d = par
    head = ratio_ab(a, b, x + 1, n + 1) / likelihood(par, x, tx, n)
    total = 0.0
    for k in range(1, horizon + 1):
        total += head * alive_through(g, d, n + k)
    return total


def expected_next_closed(par, x, tx, n, horizon):
    """Equation (13) as printed, with gamma functions: an independent route to the same number."""
    a, b, g, d = par
    head = ratio_ab(a, b, x + 1, n + 1) / likelihood(par, x, tx, n)
    lead = d / (g - 1.0) * math.exp(math.lgamma(g + d) - math.lgamma(1.0 + d))
    bracket = (math.exp(math.lgamma(1.0 + d + n) - math.lgamma(g + d + n))
               - math.exp(math.lgamma(1.0 + d + n + horizon) - math.lgamma(g + d + n + horizon)))
    return head * lead * bracket


def p_alive(par, x, tx, n):
    """Equation (11): probability of being alive at opportunity n + 1."""
    a, b, g, d = par
    return ratio_ab(a, b, x, n) * alive_through(g, d, n + 1) / likelihood(par, x, tx, n)


def cond_pmf(par, x, tx, n, horizon, xs):
    """Equation (12): probability of xs purchases over the next `horizon` opportunities."""
    a, b, g, d = par
    lik = likelihood(par, x, tx, n)
    a2 = math.comb(horizon, xs) * ratio_ab(a, b, x + xs, n + horizon) * alive_through(g, d, n + horizon)
    for i in range(xs, horizon):
        a2 += math.comb(i, xs) * ratio_ab(a, b, x + xs, n + i) * dies_at(g, d, n + i)
    out = a2 / lik
    if xs == 0:
        out += 1.0 - ratio_ab(a, b, x, n) * alive_through(g, d, n) / lik
    return out


def dert(par, x, tx, n, rate, terms=TERMS):
    """Discounted expected residual transactions: sum over k >= 1 of E[Y_(n+k)] / (1 + rate)^k."""
    a, b, g, d = par
    head = ratio_ab(a, b, x + 1, n + 1) / likelihood(par, x, tx, n)
    term = alive_through(g, d, n)
    total = 0.0
    for k in range(1, terms + 1):
        term *= (d + n + k - 1.0) / (g + d + n + k - 1.0) / (1.0 + rate)
        total += term
    return head * total


def dert_hypergeometric(par, x, tx, n, rate):
    """Equation (14) of the paper, with 2F1 evaluated by the term recursion of the authors' note."""
    a, b, g, d = par
    z = 1.0 / (1.0 + rate)
    u = 1.0
    f21 = 1.0
    for j in range(1, TERMS + 1):
        u *= (1.0 + j - 1.0) * (d + n + 1.0 + j - 1.0) / ((g + d + n + 1.0 + j - 1.0) * j) * z
        f21 += u
    return (ratio_ab(a, b, x + 1, n + 1) * alive_through(g, d, n + 1) / (1.0 + rate) * f21
            / likelihood(par, x, tx, n))


def pmf(par, n, x):
    """Equation (7): probability of x purchases in the first n opportunities."""
    a, b, g, d = par
    out = math.comb(n, x) * ratio_ab(a, b, x, n) * alive_through(g, d, n)
    for i in range(x, n):
        out += math.comb(i, x) * ratio_ab(a, b, x, i) * dies_at(g, d, i)
    return out


def mean_closed(par, n):
    """Equation (8)."""
    a, b, g, d = par
    return (a / (a + b) * d / (g - 1.0)
            * (1.0 - math.exp(math.lgamma(g + d) - math.lgamma(g + d + n) + math.lgamma(1.0 + d + n) - math.lgamma(1.0 + d))))


def pattern_probabilities(par, n):
    """P(x, t_x) = L(x, t_x) times the number of purchase strings with that recency and frequency."""
    out = []
    for x, tx in patterns(n):
        strings = 1 if x == 0 else math.comb(tx - 1, x - 1)
        out.append(strings * likelihood(par, x, tx, n))
    return out


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 check_invariants():
    """Identities the reported values rest on; the programs stop if one fails."""
    par = [TRUE_ALPHA, TRUE_BETA, TRUE_GAMMA, TRUE_DELTA]
    n = CALIBRATION
    probs = pattern_probabilities(par, n)
    by_x = [0.0] * (n + 1)
    for (x, _), pr in zip(patterns(n), probs):
        by_x[x] += pr
    checks = [
        ("pattern probabilities sum to one", plain_sum(probs), 1.0),
        ("pmf sums to one", plain_sum([pmf(par, n, x) for x in range(n + 1)]), 1.0),
        ("pmf mean equals equation (8)", plain_sum([x * pmf(par, n, x) for x in range(n + 1)]), mean_closed(par, n)),
        ("conditional pmf sums to one", plain_sum([cond_pmf(par, 2, 4, n, HOLDOUT, k) for k in range(HOLDOUT + 1)]), 1.0),
        ("conditional pmf mean equals equation (13)",
         plain_sum([k * cond_pmf(par, 2, 4, n, HOLDOUT, k) for k in range(HOLDOUT + 1)]),
         expected_next(par, 2, 4, n, HOLDOUT)),
        ("per-season terms equal the gamma form of equation (13)",
         expected_next(par, 3, 5, n, HOLDOUT), expected_next_closed(par, 3, 5, n, HOLDOUT)),
        ("DERT series equals equation (14)", dert(par, 4, 6, n, DISCOUNT), dert_hypergeometric(par, 4, 6, n, DISCOUNT)),
        ("DERT discounts the first future season once",
         dert(par, 1, 1, n, DISCOUNT, terms=1), expected_next(par, 1, 1, n, 1) / (1.0 + DISCOUNT)),
    ]
    for x in range(n + 1):
        checks.append((f"patterns with x={x} add up to the pmf", by_x[x], pmf(par, n, x)))
    for name, got, want in checks:
        if abs(got - want) > 1e-9 * max(1.0, abs(want)):
            raise SystemExit(f"error: invariant failed, {name}: {got!r} != {want!r}")
    print(f"metric.check.invariants={len(checks)} passed")


PUBLISHED = {
    # Fader and Hardie (2011), note on the Excel implementation, p. 4: LL at the starting values 1, 1, 1, 1.
    "start": "loglik=-37232.0",
    # Fader, Hardie and Shang (2010), Table 4.
    "bgbb": "alpha=1.204 beta=0.750 gamma=0.657 delta=2.783 loglik=-33225.6",
    "bb": "alpha=0.487 beta=0.826 loglik=-35516.1",
    # Table 5, row by row (x = 0, then x = 1 with t_x = 1..6, ..., x = 6), two decimals.
    "table5": "0.07 | 0.09 0.31 0.59 0.84 1.02 1.15 | 0.12 0.54 1.06 1.44 1.67 | 0.22 1.03 1.80 2.19 | 0.58 2.03 2.71 | 1.81 3.23 | 3.75",
    # Table 6, P(alive in 2002), same layout.
    "table6": "0.11 | 0.07 0.25 0.48 0.68 0.83 0.93 | 0.07 0.30 0.59 0.80 0.93 | 0.10 0.44 0.77 0.93 | 0.20 0.70 0.93 | 0.52 0.93 | 0.93",
}


def table_by_row(n, value):
    rows = [fmt(value(0, 0), 2)]
    for x in range(1, n + 1):
        rows.append(" ".join(fmt(value(x, tx), 2) for tx in range(x, n + 1)))
    return " | ".join(rows)


def published_checks():
    """Reproduce what the sources print before trusting the code on new data; stop if one is missed."""
    n = 6
    donors = plain_sum([f for _, _, f in DONORS])
    repeat = plain_sum([x * f for x, _, f in DONORS])
    if donors != 11104 or repeat != 24615:
        raise SystemExit(f"error: Table 2 transcription: {donors} donors and {repeat} repeat donations")
    print(f"metric.check.table2 patterns={len(DONORS)} donors={int(donors)} repeat_donations={int(repeat)}")
    line = f"loglik={fmt(loglik([1.0, 1.0, 1.0, 1.0], DONORS, n), 1)}"
    if line != PUBLISHED["start"]:
        raise SystemExit(f"error: Excel note starting log-likelihood not reproduced: {line}")
    print(f"metric.check.fader_hardie_2011_start {line}")
    par, ll = fit_bgbb(DONORS, n)
    line = (f"alpha={fmt(par[0], 3)} beta={fmt(par[1], 3)} gamma={fmt(par[2], 3)} delta={fmt(par[3], 3)} "
            f"loglik={fmt(ll, 1)}")
    if line != PUBLISHED["bgbb"]:
        raise SystemExit(f"error: Fader, Hardie and Shang (2010) Table 4 BG/BB not reproduced: {line}")
    print(f"metric.check.fhs_2010_table4_bgbb {line}")
    other, ll_other = fit_bgbb(DONORS, n, start=(math.log(0.01),) * 4)
    if abs(ll_other - ll) > 1e-6 or max(abs(other[i] / par[i] - 1.0) for i in range(4)) > 1e-4:
        raise SystemExit("error: the BG/BB fit depends on its starting values")
    print("metric.check.fhs_2010_table4_bgbb_from_0.01 same_optimum=yes")
    a, b, ll_bb = fit_bb(DONORS, n)
    line = f"alpha={fmt(a, 3)} beta={fmt(b, 3)} loglik={fmt(ll_bb, 1)}"
    if line != PUBLISHED["bb"]:
        raise SystemExit(f"error: Fader, Hardie and Shang (2010) Table 4 BB not reproduced: {line}")
    print(f"metric.check.fhs_2010_table4_bb {line}")
    for name, value in (("table5", lambda x, tx: expected_next(par, x, tx, n, 5)),
                        ("table6", lambda x, tx: p_alive(par, x, tx, n))):
        line = table_by_row(n, value)
        if line != PUBLISHED[name]:
            raise SystemExit(f"error: Fader, Hardie and Shang (2010) {name} not reproduced: {line}")
        print(f"metric.check.fhs_2010_{name}={line}")
    zero = 3464 * expected_next(par, 0, 0, n, 5)
    print(f"metric.check.fhs_2010_zero_repeat_donors_2002_2006={fmt(zero, 1)}")
    print(f"metric.check.fhs_2010_prior_means E(P)={fmt(par[0] / (par[0] + par[1]), 2)} "
          f"E(Theta)={fmt(par[2] / (par[2] + par[3]), 2)}")


def simulate_cohort(stream, par, count, seasons):
    """Exact draws of purchase strings by the Polya-urn form of the model: no beta variate is needed.
    At season t, an alive customer dies with probability g / (g + d + t - 1); if still alive he buys with
    probability (a + purchases so far) / (a + b + seasons alive so far)."""
    a, b, g, d = par
    out = []
    for _ in range(count):
        ys = []
        alive = True
        bought = 0
        for t in range(1, seasons + 1):
            u_die = stream.uniform()
            u_buy = stream.uniform()
            if alive and u_die < g / (g + d + t - 1.0):
                alive = False
            if alive and u_buy < (a + bought) / (a + b + t - 1.0):
                ys.append(1)
                bought += 1
            else:
                ys.append(0)
        out.append(ys)
    return out


def summarize(strings, n):
    """Recency/frequency counts over the first n opportunities, in pattern order."""
    index = {p: i for i, p in enumerate(patterns(n))}
    counts = [0] * len(index)
    for ys in strings:
        x = sum(ys[:n])
        tx = max([t + 1 for t in range(n) if ys[t] == 1], default=0)
        counts[index[(x, tx)]] += 1
    return [(x, tx, counts[i]) for i, (x, tx) in enumerate(patterns(n))]


def draw_patterns(stream, probs, count):
    """Multinomial draw of `count` customers over the patterns, by inversion of the cumulative probabilities."""
    cum = []
    total = 0.0
    for pr in probs:
        total += pr
        cum.append(total)
    counts = [0] * len(probs)
    for _ in range(count):
        u = stream.uniform() * total
        j = 0
        while j < len(cum) - 1 and u >= cum[j]:
            j += 1
        counts[j] += 1
    return counts


def naive_rate(x, n):
    return x / n


def rule_rate(x, tx, n):
    """Inactivity rule: a customer without a purchase in the last two seasons is written off."""
    return 0.0 if n - tx >= 2 else x / n


def perpetuity(rate_per_season, rate=DISCOUNT):
    """Value of a constant expected purchase rate from season n + 1 on, first season discounted once."""
    return MARGIN * rate_per_season / rate


def damped_purchases(rate_per_season, decay, horizon):
    """Damped rhythm: x/n purchases per season, times decay^k at future season k."""
    total = 0.0
    factor = 1.0
    for _ in range(horizon):
        factor *= decay
        total += factor
    return rate_per_season * total


def damped_value(rate_per_season, decay, rate=DISCOUNT):
    """Value of the damped rhythm: m (x/n) sum_k decay^k / (1 + rate)^k = m (x/n) decay / (1 + rate - decay)."""
    return MARGIN * rate_per_season * decay / (1.0 + rate - decay)


def chi2_upper_tail(stat, df):
    """P(chi2_df > stat) = 1 - P(df/2, stat/2), the regularized lower gamma by its power series."""
    a = df / 2.0
    x = stat / 2.0
    term = 1.0 / a
    total = term
    for k in range(1, 1000):
        term *= x / (a + k)
        total += term
        if term < 1e-17 * total:
            break
    return 1.0 - math.exp(a * math.log(x) - x - math.lgamma(a)) * total


def predictive_median(par, x, tx, n, horizon):
    """Smallest k with P(X <= k | history) >= 1/2 over the next `horizon` opportunities."""
    cumulative = 0.0
    for k in range(horizon + 1):
        cumulative += cond_pmf(par, x, tx, n, horizon, k)
        if cumulative >= 0.5:
            return k
    return horizon


def main():
    stream = Declared(LCG_SEED)
    true = [TRUE_ALPHA, TRUE_BETA, TRUE_GAMMA, TRUE_DELTA]
    n = CALIBRATION
    print(f"metric.design.customers={CUSTOMERS} calibration_seasons={n} holdout_seasons={HOLDOUT}")
    print(f"metric.design.seed={LCG_SEED} bootstrap={BOOTSTRAP} replications={REPLICATIONS}")
    print(f"metric.design.true_alpha={fmt(TRUE_ALPHA, 2)} true_beta={fmt(TRUE_BETA, 2)} true_gamma={fmt(TRUE_GAMMA, 2)} "
          f"true_delta={fmt(TRUE_DELTA, 2)} true_mean_p={fmt(TRUE_ALPHA / (TRUE_ALPHA + TRUE_BETA), 4)} "
          f"true_mean_theta={fmt(TRUE_GAMMA / (TRUE_GAMMA + TRUE_DELTA), 4)}")
    print(f"metric.design.margin_per_pass={fmt(MARGIN, 2)} discount_per_season={fmt(DISCOUNT, 4)} series_terms={TERMS}")
    published_checks()
    check_invariants()

    strings = simulate_cohort(stream, true, CUSTOMERS, n + HOLDOUT)
    data = summarize(strings, n)
    holdout = [sum(ys[n:]) for ys in strings]
    print("metric.data.patterns=" + " ".join(f"{x}/{tx}:{f}" for x, tx, f in data))
    buyers_by_season = [sum(ys[t] for ys in strings) for t in range(n + HOLDOUT)]
    print("metric.data.buyers_by_season=" + " ".join(f"{t + 1}:{v}" for t, v in enumerate(buyers_by_season)))
    zero = sum(f for x, _, f in data if x == 0)
    lapsed = sum(f for x, tx, f in data if n - tx >= 2)
    print(f"metric.data.calibration_repeat_purchases={sum(x * f for x, _, f in data)} "
          f"holdout_purchases={sum(holdout)} never_returned={zero} lapsed_two_seasons_or_more={lapsed}")
    print(f"metric.data.share never_returned={fmt(100.0 * zero / CUSTOMERS, 1)}% "
          f"lapsed_two_seasons_or_more={fmt(100.0 * lapsed / CUSTOMERS, 1)}%")
    # The damped shortcut lets the past rhythm decay at the mean rate at which the number of buyers holds up, seasons 1 to n.
    decay = (buyers_by_season[n - 1] / buyers_by_season[0]) ** (1.0 / (n - 1))
    print(f"metric.data.buyer_count_ratio_per_season seasons_1_to_{n}={fmt(decay, 3)}")

    par, ll = fit_bgbb(data, n)
    other, ll_other = fit_bgbb(data, n, start=(math.log(0.01),) * 4)
    if abs(ll_other - ll) > 1e-6 or max(abs(other[i] / par[i] - 1.0) for i in range(4)) > 1e-4:
        raise SystemExit("error: the synthetic BG/BB fit depends on its starting values")
    print("metric.check.synthetic_fit_from_0.01 same_optimum=yes")
    a_bb, b_bb, ll_bb = fit_bb(data, n)
    print(f"metric.fit.bgbb alpha={fmt(par[0], 3)} beta={fmt(par[1], 3)} gamma={fmt(par[2], 3)} delta={fmt(par[3], 3)} "
          f"loglik={fmt(ll, 1)}")
    print(f"metric.fit.bgbb mean_p={fmt(par[0] / (par[0] + par[1]), 4)} mean_theta={fmt(par[2] / (par[2] + par[3]), 4)}")
    print(f"metric.fit.bb alpha={fmt(a_bb, 3)} beta={fmt(b_bb, 3)} loglik={fmt(ll_bb, 1)} "
          f"likelihood_ratio={fmt(2.0 * (ll - ll_bb), 1)}")
    expected_counts = [CUSTOMERS * pmf(par, n, x) for x in range(n + 1)]
    actual_counts = [sum(f for xx, _, f in data if xx == x) for x in range(n + 1)]
    chi2 = plain_sum([(actual_counts[x] - expected_counts[x]) ** 2 / expected_counts[x] for x in range(n + 1)])
    print("metric.fit.frequency_actual_vs_bgbb=" + " ".join(
        f"{x}:{actual_counts[x]}/{fmt(expected_counts[x], 1)}" for x in range(n + 1)) + f" chi2={fmt(chi2, 2)}")
    # Fit on the 22 recency/frequency cells the model is estimated on: 22 - 1 - 4 parameters = 17 degrees of freedom.
    cell_expected = [CUSTOMERS * pr for pr in pattern_probabilities(par, n)]
    chi2_cells = plain_sum([(f - e) ** 2 / e for (_, _, f), e in zip(data, cell_expected)])
    df_cells = len(data) - 1 - 4
    print(f"metric.fit.chi2_22_cells chi2={fmt(chi2_cells, 2)} df={df_cells} p={fmt(chi2_upper_tail(chi2_cells, df_cells), 2)} "
          f"min_expected={fmt(min(cell_expected), 1)}")

    # Holdout forecasts per customer: five methods and the expectation under the true parameters.
    methods = ("naive", "rule", "damped", "bb", "bgbb", "truth")
    per_pattern = {}
    for x, tx, _ in data:
        per_pattern[(x, tx)] = {
            "naive": HOLDOUT * naive_rate(x, n),
            "rule": HOLDOUT * rule_rate(x, tx, n),
            "damped": damped_purchases(naive_rate(x, n), decay, HOLDOUT),
            "bb": HOLDOUT * (a_bb + x) / (a_bb + b_bb + n),
            "bgbb": expected_next(par, x, tx, n, HOLDOUT),
            "truth": expected_next(true, x, tx, n, HOLDOUT),
        }
    keys = []
    for ys in strings:
        x = sum(ys[:n])
        keys.append((x, max([t + 1 for t in range(n) if ys[t] == 1], default=0)))
    totals = {m: plain_sum([per_pattern[k][m] for k in keys]) for m in methods}
    print(f"metric.holdout.total actual={sum(holdout)} " + " ".join(
        f"{m}={fmt(totals[m], 1)}" for m in methods))
    print("metric.holdout.total_gap_vs_actual " + " ".join(
        f"{m}={fmt(100.0 * (totals[m] / sum(holdout) - 1.0), 1)}%" for m in methods))
    # Chance alone: standard deviation of the realized total given the histories, under the true parameters.
    variance = 0.0
    for k in keys:
        mean_k = 0.0
        second = 0.0
        for xs in range(HOLDOUT + 1):
            pr = cond_pmf(true, k[0], k[1], n, HOLDOUT, xs)
            mean_k += xs * pr
            second += xs * xs * pr
        variance += second - mean_k * mean_k
    print(f"metric.holdout.chance_sd_of_total truth_sd={fmt(math.sqrt(variance), 1)} "
          f"share_of_truth={fmt(100.0 * math.sqrt(variance) / totals['truth'], 1)}%")
    for label, group in (("frequency", lambda k: k[0]), ("recency", lambda k: k[1])):
        for v in range(n + 1):
            members = [i for i, k in enumerate(keys) if group(k) == v]
            if not members:
                continue
            m_actual = sum(holdout[i] for i in members) / len(members)
            print(f"metric.holdout.by_{label}={v} customers={len(members)} actual={fmt(m_actual, 2)} " + " ".join(
                f"{m}={fmt(plain_sum([per_pattern[keys[i]][m] for i in members]) / len(members), 2)}" for m in methods))
    mae = {m: plain_sum([abs(per_pattern[keys[i]][m] - holdout[i]) for i in range(CUSTOMERS)]) / CUSTOMERS for m in methods}
    print("metric.holdout.mean_abs_error_per_customer " + " ".join(f"{m}={fmt(mae[m], 3)}" for m in methods))
    medians = {(x, tx): predictive_median(par, x, tx, n, HOLDOUT) for x, tx, _ in data}
    mae_median = plain_sum([abs(medians[keys[i]] - holdout[i]) for i in range(CUSTOMERS)]) / CUSTOMERS
    print(f"metric.holdout.mean_abs_error_bgbb_predictive_median={fmt(mae_median, 3)}")
    active_pred = plain_sum([1.0 - cond_pmf(par, k[0], k[1], n, HOLDOUT, 0) for k in keys])
    active_true = plain_sum([1.0 - cond_pmf(true, k[0], k[1], n, HOLDOUT, 0) for k in keys])
    print(f"metric.holdout.active_customers actual={sum(1 for h in holdout if h > 0)} bgbb={fmt(active_pred, 1)} "
          f"truth={fmt(active_true, 1)}")

    for x, tx in ((6, 6), (5, 5), (4, 6), (3, 3), (1, 6), (0, 0)):
        pp = per_pattern[(x, tx)]
        print(f"metric.profile.x={x} tx={tx} customers={sum(f for xx, tt, f in data if (xx, tt) == (x, tx))} "
              f"naive={fmt(pp['naive'], 2)} rule={fmt(pp['rule'], 2)} bb={fmt(pp['bb'], 2)} "
              f"bgbb={fmt(pp['bgbb'], 2)} truth={fmt(pp['truth'], 2)} "
              f"p_alive_bgbb={fmt(p_alive(par, x, tx, n), 2)} p_alive_truth={fmt(p_alive(true, x, tx, n), 2)} "
              f"value_bgbb={fmt(MARGIN * dert(par, x, tx, n, DISCOUNT), 2)} "
              f"value_truth={fmt(MARGIN * dert(true, x, tx, n, DISCOUNT), 2)} "
              f"value_naive={fmt(perpetuity(naive_rate(x, n)), 2)}")

    # Values of the cohort: residual lifetime value, first future season discounted once.
    def cohort_value(fn):
        return plain_sum([f * fn(x, tx) for x, tx, f in data])

    values = {
        "naive": cohort_value(lambda x, tx: perpetuity(naive_rate(x, n))),
        "rule": cohort_value(lambda x, tx: perpetuity(rule_rate(x, tx, n))),
        "damped": cohort_value(lambda x, tx: damped_value(naive_rate(x, n), decay)),
        "bb": cohort_value(lambda x, tx: perpetuity((a_bb + x) / (a_bb + b_bb + n))),
        "bgbb": cohort_value(lambda x, tx: MARGIN * dert(par, x, tx, n, DISCOUNT)),
        "truth": cohort_value(lambda x, tx: MARGIN * dert(true, x, tx, n, DISCOUNT)),
    }
    print("metric.value.cohort " + " ".join(f"{m}={fmt(values[m], 0)}" for m in methods))
    print("metric.value.cohort_gap_vs_truth " + " ".join(
        f"{m}={fmt(100.0 * (values[m] / values['truth'] - 1.0), 1)}%" for m in methods))
    print("metric.value.per_customer " + " ".join(f"{m}={fmt(values[m] / CUSTOMERS, 3)}" for m in methods))
    lapsed_value = plain_sum([f * MARGIN * dert(par, x, tx, n, DISCOUNT) for x, tx, f in data if n - tx >= 2])
    lapsed_holdout = sum(holdout[i] for i, k in enumerate(keys) if n - k[1] >= 2)
    lapsed_active = sum(1 for i, k in enumerate(keys) if n - k[1] >= 2 and holdout[i] > 0)
    lapsed_expected = plain_sum([per_pattern[k]["bgbb"] for k in keys if n - k[1] >= 2])
    print(f"metric.value.written_off_by_rule customers={lapsed} bgbb_value={fmt(lapsed_value, 0)} "
          f"share_of_bgbb_value={fmt(100.0 * lapsed_value / values['bgbb'], 1)}% "
          f"holdout_purchases_actual={lapsed_holdout} holdout_purchases_bgbb={fmt(lapsed_expected, 1)} "
          f"customers_back_in_holdout={lapsed_active}")
    zero_value = plain_sum([f * MARGIN * dert(par, x, tx, n, DISCOUNT) for x, tx, f in data if x == 0])
    print(f"metric.value.never_returned customers={zero} value_each={fmt(MARGIN * dert(par, 0, 0, n, DISCOUNT), 2)} "
          f"value_total={fmt(zero_value, 0)} share_of_bgbb_value={fmt(100.0 * zero_value / values['bgbb'], 1)}%")

    def cohort_within(seasons):
        return cohort_value(lambda x, tx: MARGIN * dert(par, x, tx, n, DISCOUNT, terms=seasons))

    print("metric.horizon.share_of_bgbb_value_beyond_season=" + " ".join(
        f"{n + h}:{fmt(100.0 * (1.0 - cohort_within(h) / values['bgbb']), 1)}%" for h in (5, 10, 20)))
    for rate in (0.05, 0.10, 0.15):
        v_b = cohort_value(lambda x, tx: MARGIN * dert(par, x, tx, n, rate))
        v_t = cohort_value(lambda x, tx: MARGIN * dert(true, x, tx, n, rate))
        v_n = cohort_value(lambda x, tx: perpetuity(naive_rate(x, n), rate))
        v_r = cohort_value(lambda x, tx: perpetuity(rule_rate(x, tx, n), rate))
        v_d = cohort_value(lambda x, tx: damped_value(naive_rate(x, n), decay, rate))
        print(f"metric.sensitivity.discount={fmt(100.0 * rate, 0)}% cohort_bgbb={fmt(v_b, 0)} cohort_truth={fmt(v_t, 0)} "
              f"cohort_naive={fmt(v_n, 0)} cohort_rule={fmt(v_r, 0)} cohort_damped={fmt(v_d, 0)} "
              f"best_customer_bgbb={fmt(MARGIN * dert(par, 6, 6, n, rate), 2)}")

    probs = pattern_probabilities(par, n)
    boot_holdout = []
    boot_value = []
    boot_best = []
    boot_mean_p = []
    boot_mean_theta = []
    for _ in range(BOOTSTRAP):
        counts = draw_patterns(stream, probs, CUSTOMERS)
        b_data = [(x, tx, c) for (x, tx), c in zip(patterns(n), counts)]
        b_par, _ = fit_bgbb(b_data, n, start=[math.log(v) for v in par])
        boot_holdout.append(cohort_value(lambda x, tx: expected_next(b_par, x, tx, n, HOLDOUT)))
        boot_value.append(cohort_value(lambda x, tx: MARGIN * dert(b_par, x, tx, n, DISCOUNT)))
        boot_best.append(MARGIN * dert(b_par, 6, 6, n, DISCOUNT))
        boot_mean_p.append(b_par[0] / (b_par[0] + b_par[1]))
        boot_mean_theta.append(b_par[2] / (b_par[2] + b_par[3]))
    for v in (boot_holdout, boot_value, boot_best, boot_mean_p, boot_mean_theta):
        v.sort()
    lo = BOOTSTRAP * 25 // 1000 - 1
    hi = BOOTSTRAP * 975 // 1000 - 1
    print(f"metric.uncertainty.parametric_bootstrap draws={BOOTSTRAP} sorted_ranks={lo + 1}..{hi + 1} "
          f"holdout_total_95={fmt(boot_holdout[lo], 1)}..{fmt(boot_holdout[hi], 1)} "
          f"cohort_value_95={fmt(boot_value[lo], 0)}..{fmt(boot_value[hi], 0)} "
          f"best_customer_value_95={fmt(boot_best[lo], 2)}..{fmt(boot_best[hi], 2)}")
    print(f"metric.uncertainty.parametric_bootstrap mean_p_95={fmt(boot_mean_p[lo], 4)}..{fmt(boot_mean_p[hi], 4)} "
          f"mean_theta_95={fmt(boot_mean_theta[lo], 4)}..{fmt(boot_mean_theta[hi], 4)}")

    true_probs = pattern_probabilities(true, n)
    # The damped shortcut needs season-by-season buyers, which a recency/frequency draw does not give: it is
    # judged on the detailed cohort above, not in the replications.
    mc_methods = ("naive", "rule", "bb", "bgbb", "truth")
    rep = {m: [] for m in mc_methods}
    naive_above = 0
    rule_below = 0
    for _ in range(REPLICATIONS):
        counts = draw_patterns(stream, true_probs, CUSTOMERS)
        r_data = [(x, tx, c) for (x, tx), c in zip(patterns(n), counts)]
        r_par, _ = fit_bgbb(r_data, n)
        ra, rb, _ = fit_bb(r_data, n)
        r = {
            "naive": plain_sum([f * (HOLDOUT * naive_rate(x, n)) for x, tx, f in r_data]),
            "rule": plain_sum([f * (HOLDOUT * rule_rate(x, tx, n)) for x, tx, f in r_data]),
            "bb": plain_sum([f * (HOLDOUT * (ra + x) / (ra + rb + n)) for x, tx, f in r_data]),
            "bgbb": plain_sum([f * expected_next(r_par, x, tx, n, HOLDOUT) for x, tx, f in r_data]),
            "truth": plain_sum([f * expected_next(true, x, tx, n, HOLDOUT) for x, tx, f in r_data]),
        }
        for m in mc_methods:
            rep[m].append(r[m])
        if r["naive"] > r["truth"]:
            naive_above += 1
        if r["rule"] < r["truth"]:
            rule_below += 1
    for m in ("naive", "rule", "bb", "bgbb"):
        gaps = [rep[m][i] / rep["truth"][i] - 1.0 for i in range(REPLICATIONS)]
        mean = plain_sum(gaps) / REPLICATIONS
        sd = math.sqrt(plain_sum([(g - mean) * (g - mean) for g in gaps]) / (REPLICATIONS - 1))
        print(f"metric.monte_carlo.holdout_total_gap_vs_truth {m} mean={fmt(100.0 * mean, 1)}% "
              f"sd={fmt(100.0 * sd, 1)}% mcse={fmt(100.0 * sd / math.sqrt(REPLICATIONS), 1)}% "
              f"min={fmt(100.0 * min(gaps), 1)}% max={fmt(100.0 * max(gaps), 1)}%")
    print(f"metric.monte_carlo.naive_above_truth={naive_above} rule_below_truth={rule_below} of {REPLICATIONS}")


if __name__ == "__main__":
    main()
