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()