/*
 * reservoir.c -- experiments for the reservoir sampling article.
 *
 * Build:  gcc -O2 -Wall -Wextra -o reservoir reservoir.c -lm
 * Usage:  ./reservoir selftest
 *         ./reservoir chi SEED [T]    chi-square uniformity tests
 *                                     (T marginal trials, default 200000)
 *         ./reservoir count SEED      variate counts: Algorithm R vs L vs Z
 *         ./reservoir weighted SEED   A-Res / A-ExpJ / A-Chao inclusion tests
 *         ./reservoir merge SEED OUT  per-item inclusion frequency of merges
 *         ./reservoir trace SEED      Algorithm L jump trace (k=5)
 *
 * Every random variate drawn by the samplers goes through u01() or below(),
 * which increment `variates`.  That counter is what "random numbers used"
 * means everywhere in the article.
 */
#include <inttypes.h>
#include <math.h>
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

/* ------------------------------------------------------------ RNG ---- */
static uint64_t rs[4];
static uint64_t variates;
static uint64_t insertions;                   /* reservoir replacements */

static uint64_t splitmix64(uint64_t *x)
{
    uint64_t z = (*x += 0x9e3779b97f4a7c15ULL);
    z = (z ^ (z >> 30)) * 0xbf58476d1ce4e5b9ULL;
    z = (z ^ (z >> 27)) * 0x94d049bb133111ebULL;
    return z ^ (z >> 31);
}

static void seed_rng(uint64_t seed)
{
    for (int i = 0; i < 4; i++)
        rs[i] = splitmix64(&seed);
}

static inline uint64_t rotl(uint64_t x, int k) { return (x << k) | (x >> (64 - k)); }

static inline uint64_t next64(void) /* xoshiro256** */
{
    uint64_t r = rotl(rs[1] * 5, 7) * 9, t = rs[1] << 17;
    rs[2] ^= rs[0]; rs[3] ^= rs[1]; rs[1] ^= rs[2]; rs[0] ^= rs[3];
    rs[2] ^= t; rs[3] = rotl(rs[3], 45);
    return r;
}

/* uniform on the open interval (0, 1), 53-bit resolution */
static inline double u01(void)
{
    variates++;
    return ((double)(next64() >> 11) + 0.5) * 0x1.0p-53;
}

/* uniform integer in [0, m), m >= 1; Lemire's unbiased multiply-shift */
static inline uint64_t below(uint64_t m)
{
    variates++;
    uint64_t x = next64();
    __uint128_t p = (__uint128_t)x * m;
    uint64_t lo = (uint64_t)p;
    if (lo < m) {
        uint64_t t = -m % m;
        while (lo < t) {
            x = next64();
            p = (__uint128_t)x * m;
            lo = (uint64_t)p;
        }
    }
    return (uint64_t)(p >> 64);
}

/* ------------------------------------------------ chi-square p-value -- */
/* regularized upper incomplete gamma Q(a, x): series / continued fraction */
static double gamma_q(double a, double x)
{
    if (x <= 0) return 1.0;
    double lg = a * log(x) - x - lgamma(a);
    if (x < a + 1) {
        double sum = 1.0 / a, term = sum;
        for (int n = 1; n < 100000; n++) {
            term *= x / (a + n);
            sum += term;
            if (fabs(term) < fabs(sum) * 1e-15) break;
        }
        return 1.0 - sum * exp(lg);
    }
    double b = x + 1 - a, c = 1e300, d = 1 / b, h = d;
    for (int i = 1; i < 100000; i++) {
        double an = -i * (i - a);
        b += 2;
        d = an * d + b; if (fabs(d) < 1e-300) d = 1e-300;
        c = b + an / c; if (fabs(c) < 1e-300) c = 1e-300;
        d = 1 / d;
        double del = d * c;
        h *= del;
        if (fabs(del - 1) < 1e-15) break;
    }
    return exp(lg) * h;
}

static double chi2_sf(double x, int df) { return gamma_q(df / 2.0, x / 2.0); }

/* chi-square statistic of observed counts against expected counts */
static double chi2_stat(const uint64_t *obs, const double *expd, int cells)
{
    double s = 0;
    for (int i = 0; i < cells; i++) {
        double d = (double)obs[i] - expd[i];
        s += d * d / expd[i];
    }
    return s;
}

/* ------------------------------------------------ uniform samplers --- */
/* The stream is 0, 1, ..., n-1; each sampler writes min(n, k) indices. */
typedef void (*sampler_fn)(long n, int k, long *res);

static int fill(long n, int k, long *res)
{
    long m = n < k ? n : k;
    for (long i = 0; i < m; i++) res[i] = i;
    return n > k;
}

/* Algorithm R (Waterman; Knuth 3.4.2; Vitter 1985 Section 2) */
static void alg_r(long n, int k, long *res)
{
    if (!fill(n, k, res)) return;
    for (long i = k; i < n; i++) {           /* item i is the (i+1)-th */
        uint64_t j = below((uint64_t)i + 1); /* uniform in [0, i] */
        if (j < (uint64_t)k) { res[j] = i; insertions++; }
    }
}

/* Bug: range [0, i) instead of [0, i], i.e. accept with k/i, not k/(i+1) */
static void alg_r_offbyone(long n, int k, long *res)
{
    if (!fill(n, k, res)) return;
    for (long i = k; i < n; i++) {
        uint64_t j = below((uint64_t)i);
        if (j < (uint64_t)k) res[j] = i;
    }
}

/* Algorithm L (Li 1994).  W is the largest of the k keys in the reservoir. */
static void alg_l(long n, int k, long *res)
{
    if (!fill(n, k, res)) return;
    double W = exp(log(u01()) / k);
    long i = k - 1;                          /* index of the last item seen */
    for (;;) {
        double s = floor(log(u01()) / log1p(-W));
        if (s >= (double)(n - 1 - i)) break; /* next pick would be past the end */
        i += (long)s + 1;
        res[below((uint64_t)k)] = i;
        insertions++;
        W *= exp(log(u01()) / k);
    }
}

/* Bug: the loop starts from i = k, so item k can never be selected */
static void alg_l_startk(long n, int k, long *res)
{
    if (!fill(n, k, res)) return;
    double W = exp(log(u01()) / k);
    long i = k;
    for (;;) {
        double s = floor(log(u01()) / log(1.0 - W));
        if (s > (double)(n - i)) break;
        i += (long)s + 1;
        if (i >= n) break;
        int j = (int)(u01() * k);
        if (j >= k) j = k - 1;
        res[j] = i;
        W *= exp(log(u01()) / k);
    }
}

/*
 * Vitter's Algorithm X / Z as implemented by PostgreSQL.
 * Ported from PostgreSQL REL_17_4, src/backend/utils/misc/sampling.c,
 * reservoir_get_next_S().  Portions Copyright (c) 1996-2024, PostgreSQL
 * Global Development Group (PostgreSQL License).  sampler_random_fract()
 * is replaced by u01(); the logic is unchanged.
 */
static double pg_W;

static double pg_next_S(double t, int n)
{
    double S;
    if (t <= (22.0 * n)) {                   /* Algorithm X */
        double V = u01(), quot;
        S = 0;
        t += 1;
        quot = (t - (double)n) / t;
        while (quot > V) {
            S += 1;
            t += 1;
            quot *= (t - (double)n) / t;
        }
    } else {                                 /* Algorithm Z */
        double W = pg_W, term = t - (double)n + 1;
        for (;;) {
            double numer, numer_lim, denom, U, X, lhs, rhs, y, tmp;
            U = u01();
            X = t * (W - 1.0);
            S = floor(X);
            tmp = (t + 1) / term;
            lhs = exp(log(((U * tmp * tmp) * (term + S)) / (t + X)) / n);
            rhs = (((t + X) / (term + S)) * term) / t;
            if (lhs <= rhs) { W = rhs / lhs; break; }
            y = (((U * (t + 1)) / term) * (t + S + 1)) / (t + X);
            if ((double)n < S) { denom = t; numer_lim = term + S; }
            else { denom = t - (double)n + S; numer_lim = t + 1; }
            for (numer = t + S; numer >= numer_lim; numer -= 1) {
                y *= numer / denom;
                denom -= 1;
            }
            W = exp(-log(u01()) / n);
            if (exp(log(y) / n) <= (t + X) / t) break;
        }
        pg_W = W;
    }
    return S;
}

/* acquire_sample_rows() loop, written with jumps instead of per-row checks */
static void alg_z(long n, int k, long *res)
{
    if (!fill(n, k, res)) return;
    pg_W = exp(-log(u01()) / k);
    double t = k;
    for (;;) {
        double S = pg_next_S(t, k);
        if (t + S >= (double)n) break;
        res[(int)(k * u01())] = (long)(t + S);
        insertions++;
        t += S + 1;
    }
}

/* Systematic sampling (needs k | n): every marginal is k/n, subsets are not */
static void alg_systematic(long n, int k, long *res)
{
    long step = n / k, r = (long)below((uint64_t)step);
    for (int j = 0; j < k; j++) res[j] = r + j * step;
}

/* ------------------------------------------------------- merging ----- */
/* The stream is split into A = [0, n1) and B = [n1, n); n1 = n / 10. */
#define MAXK 1024

/* Bug: a common naive merge (each B item enters w.p. nB/(nA+nB)) */
static void merge_naive(long n, int k, long *res)
{
    long n1 = n / 10, rb[MAXK];
    alg_r(n1, k, res);                        /* dst = A */
    alg_r(n - n1, k, rb);                     /* src = B, shifted below */
    for (int i = 0; i < k; i++) {
        uint64_t j = below((uint64_t)n);
        if (j < (uint64_t)(n - n1)) res[below((uint64_t)k)] = rb[i] + n1;
    }
}

/* Correct: #items from A is hypergeometric(n, n1, k); then subsample each */
static void pick_subset(long *a, int len, int m) /* first m become the pick */
{
    for (int i = 0; i < m; i++) {
        int j = i + (int)below((uint64_t)(len - i));
        long t = a[i]; a[i] = a[j]; a[j] = t;
    }
}

static void merge_hyper(long n, int k, long *res)
{
    long n1 = n / 10, ra[MAXK], rb[MAXK], a = n1, b = n - n1;
    int ka = (int)(n1 < k ? n1 : k), h = 0;
    alg_r(n1, k, ra);
    alg_r(n - n1, k, rb);
    for (int j = 0; j < k; j++) {
        if (below((uint64_t)(a + b)) < (uint64_t)a) { h++; a--; } else b--;
    }
    pick_subset(ra, ka, h);
    pick_subset(rb, k, k - h);
    for (int j = 0; j < h; j++) res[j] = ra[j];
    for (int j = 0; j < k - h; j++) res[h + j] = rb[j] + n1;
}

/* Correct: every item draws a key; each side keeps its k smallest keys */
typedef struct { double key; long idx; } kv;

static int kv_cmp(const void *x, const void *y)
{
    double a = ((const kv *)x)->key, b = ((const kv *)y)->key;
    return (a > b) - (a < b);
}

static void bottom_k(long lo, long hi, int k, kv *out) /* max-heap of size k */
{
    int sz = 0;
    for (long i = lo; i < hi; i++) {
        kv e = { u01(), i };
        if (sz < k) {
            int c = sz++;
            while (c > 0 && out[(c - 1) / 2].key < e.key) { out[c] = out[(c - 1) / 2]; c = (c - 1) / 2; }
            out[c] = e;
        } else if (e.key < out[0].key) {
            int c = 0;
            for (;;) {
                int l = 2 * c + 1, r = l + 1, m = c;
                double mk = e.key;
                if (l < k && out[l].key > mk) { m = l; mk = out[l].key; }
                if (r < k && out[r].key > mk) m = r;
                if (m == c) break;
                out[c] = out[m]; c = m;
            }
            out[c] = e;
        }
    }
}

static void merge_bottomk(long n, int k, long *res)
{
    kv all[2 * MAXK];
    long n1 = n / 10;
    bottom_k(0, n1, k, all);
    bottom_k(n1, n, k, all + k);
    qsort(all, 2 * (size_t)k, sizeof(kv), kv_cmp);
    for (int j = 0; j < k; j++) res[j] = all[j].idx;
}

/* ------------------------------------------------------ chi tests ---- */
/* Marginal counts are Binomial(T, k/n) with covariance -T p(1-p)/(n-1),
 * so (n-1)/(n-k) * Pearson is asymptotically chi-square with n-1 df. */
static void chi_marginal(const char *name, sampler_fn f, long n, int k,
                         long trials, uint64_t *cnt)
{
    long res[MAXK];
    memset(cnt, 0, sizeof(uint64_t) * (size_t)n);
    for (long t = 0; t < trials; t++) {
        f(n, k, res);
        for (int j = 0; j < k; j++) cnt[res[j]]++;
    }
    double *e = malloc(sizeof(double) * (size_t)n);
    for (long i = 0; i < n; i++) e[i] = (double)trials * k / n;
    double x = chi2_stat(cnt, e, (int)n) * (n - 1) / (double)(n - k);
    printf("marginal\t%s\t%.1f\t%ld\t%.3g\n", name, x, n - 1, chi2_sf(x, (int)n - 1));
    free(e);
}

static void chi_joint(const char *name, sampler_fn f, long trials)
{
    enum { N = 6, K = 3 };
    int id[1 << N], cells = 0;
    uint64_t cnt[20] = { 0 };
    double e[20];
    for (int m = 0; m < (1 << N); m++)
        id[m] = __builtin_popcount((unsigned)m) == K ? cells++ : -1;
    for (long t = 0; t < trials; t++) {
        long res[K];
        int m = 0;
        f(N, K, res);
        for (int j = 0; j < K; j++) m |= 1 << res[j];
        cnt[id[m]]++;
    }
    for (int c = 0; c < cells; c++) e[c] = (double)trials / cells;
    double x = chi2_stat(cnt, e, cells);
    printf("joint\t%s\t%.1f\t%d\t%.3g\n", name, x, cells - 1, chi2_sf(x, cells - 1));
}

/* ------------------------------------------------ weighted samplers -- */
static void minheap_sift(kv *h, int m, int c)
{
    kv e = h[c];
    for (;;) {
        int l = 2 * c + 1, r = l + 1, s = c;
        double sk = e.key;
        if (l < m && h[l].key < sk) { s = l; sk = h[l].key; }
        if (r < m && h[r].key < sk) s = r;
        if (s == c) break;
        h[c] = h[s]; c = s;
    }
    h[c] = e;
}

static void heap_out(kv *h, int m, long *res)
{
    for (int i = 0; i < m; i++) res[i] = h[i].idx;
}

/* A-Res (Efraimidis & Spirakis 2006): keep the m largest keys u^(1/w).
 * use_pow = 1 computes the key literally; 0 keeps log(key) = log(u)/w. */
static double ares_key(double w, int use_pow)
{
    double u = u01();
    return use_pow ? pow(u, 1.0 / w) : log(u) / w;
}

static void a_res(const double *w, long n, int m, long *res, int use_pow)
{
    kv h[MAXK];
    for (int i = 0; i < m; i++) h[i] = (kv){ ares_key(w[i], use_pow), i };
    for (int i = m / 2 - 1; i >= 0; i--) minheap_sift(h, m, i);
    for (long i = m; i < n; i++) {
        double key = ares_key(w[i], use_pow);
        if (key > h[0].key) {                 /* ties keep the older item */
            h[0] = (kv){ key, i };
            minheap_sift(h, m, 0);
            insertions++;
        }
    }
    heap_out(h, m, res);
}

/* A-ExpJ: exponential jumps over cumulative weight; keys in log domain */
static void a_expj(const double *w, long n, int m, long *res)
{
    kv h[MAXK];
    for (int i = 0; i < m; i++) h[i] = (kv){ log(u01()) / w[i], i };
    for (int i = m / 2 - 1; i >= 0; i--) minheap_sift(h, m, i);
    long i = m;
    while (i < n) {
        double logT = h[0].key;               /* log of threshold T_w */
        double X = log(u01()) / logT;         /* weight to skip */
        double acc = 0;
        while (i < n && (acc += w[i]) < X) i++;
        if (i >= n) break;
        double tw = exp(w[i] * logT);         /* T_w^{w_i} */
        double r2 = tw + u01() * (1.0 - tw);  /* uniform on (t_w, 1) */
        h[0] = (kv){ log(r2) / w[i], i };
        minheap_sift(h, m, 0);
        insertions++;
        i++;
    }
    heap_out(h, m, res);
}

/* A-Chao (Chao 1982), only for streams where m*w_i/W_i <= 1 always holds
 * and the first m weights are equal, so no overweight handling is needed. */
static void a_chao(const double *w, long n, int m, long *res)
{
    double W = 0;
    for (int i = 0; i < m; i++) { res[i] = i; W += w[i]; }
    for (long i = m; i < n; i++) {
        W += w[i];
        double p = m * w[i] / W;
        if (p > 1) { fprintf(stderr, "a_chao: infeasible weight\n"); exit(1); }
        if (u01() < p) { res[below((uint64_t)m)] = i; insertions++; }
    }
}

/* ------------------------------------------------------- commands ---- */
static double harmonic(double n) /* H_n, asymptotic for large n */
{
    if (n < 1000) { double h = 0; for (long i = 1; i <= (long)n; i++) h += 1.0 / i; return h; }
    return log(n) + 0.57721566490153286 + 1 / (2 * n) - 1 / (12 * n * n);
}

static int check_sample(const long *res, long n, int k)
{
    for (int i = 0; i < k; i++) {
        if (res[i] < 0 || res[i] >= n) return 0;
        for (int j = 0; j < i; j++) if (res[i] == res[j]) return 0;
    }
    return 1;
}

static int cmd_selftest(void)
{
    struct { double x; int df; } q[] = { { 3.841459, 1 }, { 30.14353, 19 }, { 1073.643, 999 } };
    int ok = 1;
    for (int i = 0; i < 3; i++) {
        double p = chi2_sf(q[i].x, q[i].df);
        printf("chi2_sf(%g, %d) = %.5f (expect 0.05)\n", q[i].x, q[i].df, p);
        ok &= fabs(p - 0.05) < 1e-3;
    }
    sampler_fn f[] = { alg_r, alg_l, alg_z, merge_hyper, merge_bottomk };
    long res[MAXK];
    seed_rng(1);
    for (int s = 0; s < 5; s++)
        for (long n = 1; n <= 3000; n += 7) {
            int k = 1 + (int)(n % 50);
            if ((s == 3 || s == 4) && n / 10 < k) continue;
            f[s](n, k, res);
            ok &= check_sample(res, n, (int)(n < k ? n : k));
        }
    printf("selftest %s\n", ok ? "PASS" : "FAIL");
    return !ok;
}

static void cmd_chi(long trials)
{
    static uint64_t cnt[1000];
    struct { const char *name; sampler_fn f; } M[] = {
        { "R", alg_r }, { "L", alg_l }, { "Z(PostgreSQL)", alg_z },
        { "merge-hypergeometric", merge_hyper }, { "merge-bottom-k", merge_bottomk },
        { "BUG:R-range-[0,i)", alg_r_offbyone }, { "BUG:L-start-at-k", alg_l_startk },
        { "BUG:merge-naive", merge_naive },
    }, J[] = {
        { "R", alg_r }, { "L", alg_l }, { "Z(PostgreSQL)", alg_z },
        { "systematic", alg_systematic },
        { "BUG:R-range-[0,i)", alg_r_offbyone }, { "BUG:L-start-at-k", alg_l_startk },
    };
    printf("# test\tsampler\tstatistic\tdf\tp\n");
    printf("# marginal: n=1000 k=10 trials=%ld; joint: n=6 k=3 trials=1000000\n", trials);
    for (size_t i = 0; i < sizeof M / sizeof M[0]; i++)
        chi_marginal(M[i].name, M[i].f, 1000, 10, trials, cnt);
    for (size_t i = 0; i < sizeof J / sizeof J[0]; i++)
        chi_joint(J[i].name, J[i].f, 1000000);
}

static void cmd_count(void)
{
    static long res[MAXK];
    int ks[] = { 10, 100, 1000 };
    printf("# k\tn\talgorithm\tvariates\treplacements\tE[replacements]=k(H_n-H_k)\n");
    for (int a = 0; a < 3; a++)
        for (long n = 1000; n <= 100000000L; n *= 10) {
            int k = ks[a];
            if (n <= k) continue;
            double er = k * (harmonic((double)n) - harmonic(k));
            struct { const char *name; sampler_fn f; } A[] = { { "R", alg_r }, { "L", alg_l }, { "Z", alg_z } };
            for (int i = 0; i < 3; i++) {
                variates = insertions = 0;
                A[i].f(n, k, res);
                printf("%d\t%ld\t%s\t%" PRIu64 "\t%" PRIu64 "\t%.1f\n", k, n, A[i].name,
                       variates, insertions, er);
            }
        }
}

/* exact pair distribution of A-Chao on a tiny stream (m = 2) */
static void chao_rec(const double *w, int n, int i, double W, long *sl,
                     double pr, double *out)
{
    if (i == n) {
        long a = sl[0] < sl[1] ? sl[0] : sl[1], b = sl[0] ^ sl[1] ^ a;
        out[a * n + b] += pr;
        return;
    }
    W += w[i];
    double p = 2 * w[i] / W;
    chao_rec(w, n, i + 1, W, sl, pr * (1 - p), out);
    for (int s = 0; s < 2; s++) {
        long keep = sl[s];
        sl[s] = i;
        chao_rec(w, n, i + 1, W, sl, pr * p / 2, out);
        sl[s] = keep;
    }
}

enum { WS_RES_LOG, WS_RES_POW, WS_EXPJ, WS_CHAO };

static void run_weighted(const char *name, int alg, const double *w,
                         const double *exact, long trials)
{
    uint64_t pc[16] = { 0 }, inc[4] = { 0 }, obs[6];
    double e[6];
    long res[2];
    for (long t = 0; t < trials; t++) {
        if (alg == WS_RES_LOG) a_res(w, 4, 2, res, 0);
        else if (alg == WS_RES_POW) a_res(w, 4, 2, res, 1);
        else if (alg == WS_EXPJ) a_expj(w, 4, 2, res);
        else a_chao(w, 4, 2, res);
        long a = res[0] < res[1] ? res[0] : res[1], b = res[0] ^ res[1] ^ a;
        pc[a * 4 + b]++;
        inc[res[0]]++; inc[res[1]]++;
    }
    int c = 0;
    for (int a = 0; a < 4; a++)
        for (int b = a + 1; b < 4; b++, c++) { obs[c] = pc[a * 4 + b]; e[c] = exact[a * 4 + b] * trials; }
    double x = chi2_stat(obs, e, 6);
    printf("weighted\t%s\tincl=%.4f,%.4f,%.4f,%.4f\tchi2=%.1f\tdf=5\tp=%.3g\n", name,
           (double)inc[0] / trials, (double)inc[1] / trials, (double)inc[2] / trials,
           (double)inc[3] / trials, x, chi2_sf(x, 5));
}

static void cmd_weighted(void)
{
    const long T = 1000000;
    double w[4] = { 1, 1, 1, 2 }, nw[16] = { 0 }, chao[16] = { 0 }, uni[16] = { 0 };
    double tot = 5;
    for (int a = 0; a < 4; a++)
        for (int b = a + 1; b < 4; b++) {
            nw[a * 4 + b] = w[a] / tot * w[b] / (tot - w[a]) + w[b] / tot * w[a] / (tot - w[b]);
            uni[a * 4 + b] = 1.0 / 6;
        }
    long sl[2] = { 0, 1 };
    chao_rec(w, 4, 2, 2, sl, 1.0, chao);
    printf("# weights 1,1,1,2  m=2  trials=%ld; exact WRS-N-W incl = 0.4333 x3, 0.7000\n", T);
    run_weighted("A-Res(log-key)", WS_RES_LOG, w, nw, T);
    run_weighted("A-ExpJ", WS_EXPJ, w, nw, T);
    run_weighted("A-Chao(vs its exact law)", WS_CHAO, w, chao, T);
    run_weighted("A-Chao(vs WRS-N-W)", WS_CHAO, w, nw, T);

    double ws[4] = { 5e-4, 5e-4, 5e-4, 5e-4 };
    long zero = 0;
    for (long t = 0; t < T; t++) zero += pow(u01(), 1.0 / ws[0]) == 0.0;
    printf("# equal weights 5e-4: pow(u, 1/w) == 0 for %.4f of keys; uniform expected\n", (double)zero / T);
    run_weighted("A-Res(pow-key)", WS_RES_POW, ws, uni, T);
    run_weighted("A-Res(log-key)", WS_RES_LOG, ws, uni, T);

    const long n = 1000000;
    const int m = 100;
    double *wv = malloc(sizeof(double) * n);
    long res[MAXK];
    printf("# n=%ld m=%d; E[insertions] for exchangeable weights = m(H_n-H_m) = %.1f\n", n, m,
           m * (harmonic(n) - harmonic(m)));
    const char *pat[] = { "iid-uniform", "increasing(i+1)", "decreasing(n-i)", "geometric(1.0001^i)" };
    for (int p = 0; p < 4; p++) {
        for (long i = 0; i < n; i++)
            wv[i] = p == 0 ? u01() : p == 1 ? (double)(i + 1) : p == 2 ? (double)(n - i) : exp(i * log1p(1e-4));
        variates = insertions = 0;
        a_res(wv, n, m, res, 0);
        printf("jumps\t%s\tA-Res\tvariates=%" PRIu64 "\tinsertions=%" PRIu64 "\n", pat[p], variates, insertions);
        variates = insertions = 0;
        a_expj(wv, n, m, res);
        printf("jumps\t%s\tA-ExpJ\tvariates=%" PRIu64 "\tinsertions=%" PRIu64 "\n", pat[p], variates, insertions);
    }
    free(wv);
}

static void cmd_merge(const char *out)
{
    static uint64_t c[3][1000];
    chi_marginal("BUG:merge-naive", merge_naive, 1000, 10, 200000, c[0]);
    chi_marginal("merge-hypergeometric", merge_hyper, 1000, 10, 200000, c[1]);
    chi_marginal("merge-bottom-k", merge_bottomk, 1000, 10, 200000, c[2]);
    FILE *f = fopen(out, "w");
    if (!f) { perror(out); exit(1); }
    fprintf(f, "item\tnaive\thypergeometric\tbottomk\n");
    for (int i = 0; i < 1000; i++)
        fprintf(f, "%d\t%.6f\t%.6f\t%.6f\n", i, c[0][i] / 2e5, c[1][i] / 2e5, c[2][i] / 2e5);
    fclose(f);
}

static void cmd_trace(void)
{
    const int k = 5;
    const long n = 3000;
    double W = exp(log(u01()) / k);
    long i = k - 1;
    printf("# Algorithm L, k=%d n=%ld: event\tindex\tW after the event\n", k, n);
    printf("init\t%ld\t%.6g\n", i, W);
    for (;;) {
        double s = floor(log(u01()) / log1p(-W));
        if (s >= (double)(n - 1 - i)) break;
        i += (long)s + 1;
        (void)below((uint64_t)k);
        W *= exp(log(u01()) / k);
        printf("accept\t%ld\t%.6g\n", i, W);
    }
}

int main(int argc, char **argv)
{
    if (argc < 2) { fprintf(stderr, "usage: see header\n"); return 2; }
    if (!strcmp(argv[1], "selftest")) return cmd_selftest();
    uint64_t seed = argc > 2 ? strtoull(argv[2], NULL, 10) : 1;
    seed_rng(seed);
    printf("# seed %" PRIu64 "\n", seed);
    if (!strcmp(argv[1], "chi")) cmd_chi(argc > 3 ? atol(argv[3]) : 200000);
    else if (!strcmp(argv[1], "count")) cmd_count();
    else if (!strcmp(argv[1], "weighted")) cmd_weighted();
    else if (!strcmp(argv[1], "merge") && argc > 3) cmd_merge(argv[3]);
    else if (!strcmp(argv[1], "trace")) cmd_trace();
    else { fprintf(stderr, "unknown command\n"); return 2; }
    return 0;
}
