drf.py

11.2 kB · python · 291 lines

1import argparse2import math3import sys4import time5from fractions import Fraction67import numpy as np89A002487 = [0, 1, 1, 2, 1, 3, 2, 3, 1, 4, 3, 5, 2, 5, 3, 4, 1, 5, 4, 7, 3, 8, 5, 7, 2, 7, 5, 8, 3, 7, 4, 5]10PHI = (1 + 5 ** 0.5) / 21112def product(q, b, L):13    r = [1]14    for j in range(L):15        s = b ** j16        out = [0] * (len(r) + s * (len(q) - 1))17        for k, w in enumerate(q):18            if w:19                for n, v in enumerate(r):20                    out[n + k * s] += w * v21        r = out22    return r2324def digits(q, b, L, n, memo):25    if L == 0:26        return 1 if n == 0 else 027    key = (L, n)28    if key not in memo:29        memo[key] = sum(w * digits(q, b, L - 1, (n - c) // b, memo) for c, w in enumerate(q) if w and c <= n and (n - c) % b == 0)30    return memo[key]3132def stern(m):33    s = [0, 1]34    for n in range(2, m + 1):35        s.append(s[n // 2] if n % 2 == 0 else s[n // 2] + s[n // 2 + 1])36    return s3738def fib(n):39    a, b = 0, 140    for _ in range(n):41        a, b = b, a + b42    return a4344def count(a):45    q = [1, 1, 1]46    s = stern(2 ** (a.top + 1) + 2)47    assert s[:len(A002487)] == A00248748    print(f"count: first {len(A002487)} terms of A002487 agree with the recursion")49    bad = 050    for L in range(1, a.exact + 1):51        r = product(q, 2, L)52        memo = {}53        bad += sum(r[n] != digits(q, 2, L, n, memo) for n in range(len(r)))54    print(f"count: product formula against the digit recursion, Q = 1+z+z^2, b = 2, L = 1..{a.exact}: {bad} mismatches")55    for Q, b in [([1, 1], 2), ([2, 1], 2), ([1, 1, 1, 1], 3), ([2, 2, 3, 2, 1], 2)]:56        memo = {}57        r = product(Q, b, a.exact - 4)58        print(f"count: Q = {Q}, b = {b}, L = {a.exact - 4}: {sum(r[n] != digits(Q, b, a.exact - 4, n, memo) for n in range(len(r)))} mismatches over {len(r)} lags")59    print("L lags total mean max F_(L+1) argmax single single_law stern_mismatch mirror_mismatch fold_mismatch peak/mean ratio")60    prev = None61    for L in range(1, a.top + 1):62        r = product(q, 2, L)63        N = len(r)64        assert N == 2 ** (L + 1) - 1 and sum(r) == 3 ** L65        sm = sum(r[n] != s[n + 1] for n in range(2 ** L))66        mm = sum(r[n] != r[N - 1 - n] for n in range(N))67        top = max(r)68        arg = [n for n in range(N) if r[n] == top]69        m = (2 ** (L + 1) + (-1) ** L) // 370        law = sorted({m - 1, 3 * 2 ** (L - 1) - m - 1, N - m, N - 3 * 2 ** (L - 1) + m})71        one = [n for n in range(N) if r[n] == 1]72        onelaw = sorted({2 ** k - 1 for k in range(L + 1)} | {2 ** (L + 1) - 1 - 2 ** k for k in range(L + 1)})73        fm = sum(s[2 ** L + j] != (r[j - 1] if j else 0) + r[j - 1 + 2 ** L] for j in range(2 ** L))74        assert top == fib(L + 1) and arg == law and one == onelaw and len(one) == 2 * L + 1 and sm == 0 and mm == 0 and fm == 075        pm = Fraction(top * N, 3 ** L)76        ratio = float(pm / prev) if prev else float("nan")77        prev = pm78        shown = arg if len(arg) <= 4 else arg[:4]79        print(f"{L} {N} {3 ** L} {3 ** L / N:.4f} {top} {fib(L + 1)} {shown} {len(one)} {2 * L + 1} {sm} {mm} {fm} {float(pm):.4f} {ratio:.5f}")80    print(f"count: 2 phi/3 = {2 * PHI / 3:.5f}")8182def cyclo(n, memo={}):83    if n not in memo:84        p = [Fraction(-1)] + [Fraction(0)] * (n - 1) + [Fraction(1)]85        for d in range(1, n):86            if n % d == 0:87                p = divide(p, cyclo(d))[0]88        memo[n] = p89    return memo[n]9091def divide(p, d):92    p = list(p)93    out = [Fraction(0)] * max(1, len(p) - len(d) + 1)94    while len(p) >= len(d) and any(p):95        c = p[-1] / d[-1]96        k = len(p) - len(d)97        out[k] = c98        for i, v in enumerate(d):99            p[k + i] -= c * v100        p.pop()101    return out, p102103def vanishes(q, n):104    return not any(divide(q, cyclo(n))[1])105106def verdict(q, b, depth):107    q = [Fraction(x) for x in q]108    while q and q[-1] == 0:109        q.pop()110    for k in range(1, b ** depth):111        if k % b == 0:112            continue113        if not any(vanishes(q, b ** i // math.gcd(k, b ** i)) for i in range(1, depth + 1)):114            return False, k115    return True, None116117def fourier(q, b, t, terms=80):118    q = np.array([float(x) for x in q])119    z = 1 + 0j120    for i in range(1, terms):121        w = np.exp(2j * np.pi * t / b ** i) ** np.arange(len(q))122        z *= (q @ w) / q.sum()123    return abs(z)124125def family():126    rows = []127    for b in [2, 3, 4]:128        for K in range(2, 9):129            rows.append((f"uniform K={K}", [1] * K, b))130    for c in ["1/9", "1/4", "1/2", "1", "2", "4"]:131        cf = Fraction(c)132        rows.append((f"1+c(1+z) c={c}", [1 + cf, cf], 2))133    for c in ["1/81", "1/9", "1/4", "1/2", "1", "2", "4"]:134        cf = Fraction(c)135        rows.append((f"1+c(1+z+z^2)^2 c={c}", [1 + cf, 2 * cf, 3 * cf, 2 * cf, cf], 2))136    rows.append(("(1+z+z^2)^2", [1, 2, 3, 2, 1], 2))137    return rows138139def limit(a):140    print("Q b verdict witness_k max|hat mu(k)|,k<b^3 |hat mu(1)|")141    agree = 0142    for name, q, b in family():143        ok, k = verdict(q, b, a.depth)144        top = max(fourier(q, b, t) for t in range(1, b ** 3) if t % b)145        print(f"{name} {b} {'absolutely continuous' if ok else 'singular'} {k if k else '-'} {top:.2e} {fourier(q, b, 1):.4f}")146        assert ok == (top < 1e-9)147        agree += 1148    print(f"limit: {agree} rows, verdict agrees with the Fourier product on every row")149    s2 = 2 ** 0.5150    q = np.polymul(np.polymul([1, -s2, 1], [1, 0, s2, 0, 1]), [1, 2, 3, 2, 1])[::-1]151    at = lambda t: abs(np.polyval(q[::-1], np.exp(2j * np.pi * t)))152    cover = [min(i for i in range(1, a.depth + 1) if at(k / 2 ** i) < 1e-12) for k in range(1, 2 ** a.depth, 2)]153    off = [max(at(k / 2 ** i) for k in range(1, 2 ** i, 2)) for i in range(1, 5)]154    top = max(fourier(q, 2, t) for t in range(1, 256, 2))155    print(f"limit: real Q = (z^2 - sqrt2 z + 1)(z^4 + sqrt2 z^2 + 1)(1+z+z^2)^2, b = 2: min coefficient {q.min():.4f}, cover level of odd k < 2^{a.depth} {sorted(set(cover))}, max |Q| at primitive 2^i-th roots i = 1..4 {[round(float(x), 4) for x in off]}, max|hat mu(t)| over odd t < 256 {top:.1e}")156    for name, q, b in [("uniform K=3", [1, 1, 1], 2), ("uniform K=4", [1, 1, 1, 1], 2), ("1+(1+z+z^2)^2", [2, 2, 3, 2, 1], 2), ("1+(1+z)/4", [5, 1], 2)]:157        cells = []158        for L in [8, 12, 16, 20]:159            r = np.array(product(q, b, L), dtype=float)160            p = np.sort(r)[::-1] / r.sum()161            cells.append(f"L={L}: {np.searchsorted(np.cumsum(p), 0.5) / len(p):.4f}")162        print(f"limit: {name}, share of lags carrying half the mass, {', '.join(cells)}")163164def stack(rng, C, L, b, K, sigma, resid):165    A = rng.normal(0, sigma, (L, K, C, C))166    if resid:167        A[:, 0] += np.eye(C)168    return A169170def grad_filter(A, u, b, N):171    L, K, C, _ = A.shape172    g = np.zeros((N, C))173    g[0] = u174    for j in range(L - 1, -1, -1):175        s = b ** j176        new = np.zeros_like(g)177        for k in range(K):178            new[k * s:] += g[:N - k * s] @ A[j, k]179        g = new180    return g181182def gradient(a):183    rng = np.random.default_rng(1)184    C, L, b = a.channels, a.levels, 2185    for name, K, resid, cs in [("pure K=3", 3, False, 1.0), ("residual 1+c(1+z)", 2, True, 0.5)]:186        sigma = (cs / C) ** 0.5187        q = [1 + cs, cs] if resid else [cs] * K188        exact = np.array(product([Fraction(x).limit_denominator(1000) for x in q], b, L), dtype=float)189        N = len(exact)190        acc = np.zeros(N)191        sq = np.zeros(N)192        u = np.zeros(C)193        u[0] = 1.0194        for _ in range(a.draws):195            g = (grad_filter(stack(rng, C, L, b, K, sigma, resid), u, b, N) ** 2).sum(1)196            acc += g197            sq += g * g198        mean = acc / a.draws199        se = np.sqrt(np.maximum(sq / a.draws - mean ** 2, 0) / a.draws)200        z = np.abs(mean - exact) / np.maximum(se, 1e-300)201        rel = np.abs(mean - exact) / exact202        print(f"gradient: {name}, C={C}, L={L}, {a.draws} draws, {N} lags: max |z| {z.max():.2f}, share |z|>3 {np.mean(z > 3):.4f}, median rel err {np.median(rel):.4f}, corr {np.corrcoef(mean, exact)[0, 1]:.4f}")203204def train(A, u, v, n, lr, cap, hit):205    L, K, C, _ = A.shape206    b = 2207    N = 2 ** (L + 1) - 1208    for step in range(1, cap + 1):209        F = [np.zeros((N, C))]210        F[0][0] = v211        for j in range(L):212            s = b ** j213            new = np.zeros((N, C))214            for k in range(K):215                new[k * s:] += F[j][:N - k * s] @ A[j, k].T216            F.append(new)217        f = F[L] @ u218        if f[n] >= hit:219            return step220        e = 2 * f221        e[n] -= 2222        G = np.outer(e, u)223        gu = e @ F[L]224        gA = np.zeros_like(A)225        for j in range(L - 1, -1, -1):226            s = b ** j227            back = np.zeros((N, C))228            for k in range(K):229                gA[j, k] = G[k * s:].T @ F[j][:N - k * s]230                back[:N - k * s] += G[k * s:] @ A[j, k]231            G = back232        gv = G[0]233        A -= lr * gA234        u -= lr * gu235        v -= lr * gv236    return None237238def copy(a):239    L, C = a.levels, a.channels240    r = product([1, 1, 1], 2, L)241    lags = []242    for want in [1, 2, 3, 5, 8, 13, 21, 34, 55]:243        pick = [n for n in range(2 ** (L - 1), 2 ** L) if r[n] == want]244        if pick:245            lags.append(pick[len(pick) // 2])246    print(f"copy: K=3, b=2, L={L}, C={C}, lags {lags}, r {[r[n] for n in lags]}, seeds {a.seeds}, lr {a.lr}, hit f(n) >= {a.hit}")247    rows = []248    capped = 0249    for n in lags:250        steps = []251        for seed in range(a.seeds):252            rng = np.random.default_rng(100 + seed)253            A = rng.normal(0, a.scale * (1 / (3 * C)) ** 0.5, (L, 3, C, C))254            u = rng.normal(0, 1 / C ** 0.5, C)255            v = rng.normal(0, 1, C)256            steps.append(train(A, u, v, n, a.lr, a.cap, a.hit))257        med = float(np.median([x if x else math.inf for x in steps]))258        capped += steps.count(None)259        rows.append((n, r[n], med))260        print(f"copy: lag {n} r {r[n]} steps {steps} median {med}")261    x = np.log([1 / t[1] for t in rows])262    y = np.log([t[2] for t in rows])263    good = np.isfinite(y)264    beta, alpha = np.polyfit(x[good], y[good], 1)265    print(f"copy: {capped} runs capped at {a.cap} steps, read as above the cap in each median")266    print(f"copy: log steps against log 1/r over {good.sum()} lags: slope {beta:.3f}, corr {np.corrcoef(x[good], y[good])[0, 1]:.3f}, steps ratio r=1 over r=max {rows[0][2] / rows[-1][2]:.2f} against r ratio {rows[-1][1] / rows[0][1]}")267268def main():269    p = argparse.ArgumentParser()270    p.add_argument("verb", choices=["count", "limit", "gradient", "copy", "all"])271    p.add_argument("--top", type=int, default=20)272    p.add_argument("--exact", type=int, default=12)273    p.add_argument("--depth", type=int, default=6)274    p.add_argument("--channels", type=int, default=8)275    p.add_argument("--levels", type=int, default=8)276    p.add_argument("--draws", type=int, default=4000)277    p.add_argument("--seeds", type=int, default=5)278    p.add_argument("--lr", type=float, default=0.005)279    p.add_argument("--scale", type=float, default=1.0)280    p.add_argument("--cap", type=int, default=20000)281    p.add_argument("--hit", type=float, default=0.5)282    a = p.parse_args()283    verbs = ["count", "limit", "gradient", "copy"] if a.verb == "all" else [a.verb]284    for verb in verbs:285        t = time.time()286        globals()[verb](a)287        print(f"{verb}: {time.time() - t:.1f} seconds")288        sys.stdout.flush()289290if __name__ == "__main__":291    main()