Spectral-gap adaptive polynomial filtering / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7import sys
  8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  9from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
 10
 11SEEDS = tuple(range(8))
 12SWEEP_SEEDS = tuple(range(4))
 13EPOCHS = 12
 14BATCH = 128
 15# This is the only new solver/readout knob; it is fixed before running.
 16GAP = 0.25
 17K = 7
 18
 19
 20def fejer_coeffs(k):
 21    return np.ones(k + 1, dtype=np.float64) / (k + 1)
 22
 23
 24def jackson_coeffs(k):
 25    n = k // 2 + 1
 26    return np.convolve(np.ones(n) / n, np.ones(n) / n)
 27
 28
 29def math_checks():
 30    out = {}
 31    for k in (3, 7, 15):
 32        f, j = fejer_coeffs(k), jackson_coeffs(k)
 33        out[str(k)] = {
 34            "fejer_p1_error": float(abs(f.sum() - 1)),
 35            "jackson_p1_error": float(abs(j.sum() - 1)),
 36            "jackson_degree": int(len(j) - 1),
 37            "nonnegative": bool(j.min() >= 0),
 38        }
 39    rng = np.random.default_rng(11)
 40    a = rng.normal(size=20)
 41    # Direct polynomial versus the explicit reflection iterates.
 42    c = fejer_coeffs(15)
 43    direct = np.zeros_like(a)
 44    z = a.copy()
 45    for x in c:
 46        direct += x * z
 47        z = (2 * 0.75 - 1) * z
 48    horner = np.zeros_like(a)
 49    t = 2 * .75 - 1
 50    for x in c[::-1]:
 51        horner = horner * t + x * a
 52    out["iterate_identity_error"] = float(np.linalg.norm(direct - horner))
 53    return out
 54
 55
 56def filter_output(raw, method="idea", k=K, gap=GAP):
 57    """Apply p(2F-I) to y0=0, where F(y)=(1-gap)y+gap*raw.
 58
 59    This is an end-to-end differentiable filtered prediction. The baseline
 60    uses the same rnn_small and optimizer but returns raw predictions.
 61    """
 62    if method == "baseline":
 63        return raw
 64    # Spectral-gap selector: remain conservative at the critical scale.
 65    if gap * k < 2.0:
 66        coeff_np = fejer_coeffs(k)
 67    else:
 68        coeff_np = jackson_coeffs(k)
 69    coeff = torch.as_tensor(coeff_np, dtype=raw.dtype, device=raw.device)
 70    # Reflection map applied to a state, with raw as the fixed-point forcing.
 71    z = torch.zeros_like(raw)
 72    total = torch.zeros_like(raw)
 73    for c in coeff:
 74        total = total + c * z
 75        z = 2 * ((1 - gap) * z + gap * raw) - z
 76    # Safety check analogous to the proposal, evaluated per batch.
 77    fejer = torch.zeros_like(raw)
 78    zf = torch.zeros_like(raw)
 79    for _ in range(k + 1):
 80        fejer = fejer + zf / (k + 1)
 81        zf = 2 * ((1 - gap) * zf + gap * raw) - zf
 82    if torch.mean(total.square()) > 1.10 * torch.mean(fejer.square()):
 83        return fejer
 84    return total
 85
 86
 87def seed_all(seed):
 88    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 89    if torch.cuda.is_available():
 90        torch.cuda.manual_seed_all(seed)
 91
 92
 93def train_one(seed, lr, method):
 94    seed_all(seed)
 95    d = get_dataset("dynamics", seed, n_train=400, n_test=200)
 96    model = make_model("rnn_small", d["input_shape"], d["out_dim"])
 97    lossf = nn.MSELoss()
 98    device = "cuda" if torch.cuda.is_available() else "cpu"
 99    try:
100        model.to(device)
101        xtr, ytr = d["xtr"].to(device), d["ytr"].to(device)
102        opt = torch.optim.Adam(model.parameters(), lr=lr)
103        for _ in range(EPOCHS):
104            model.train()
105            perm = torch.randperm(len(xtr), device=device)
106            for i in range(0, len(xtr), BATCH):
107                ix = perm[i:i+BATCH]
108                raw = model(xtr[ix])
109                pred = filter_output(raw, method)
110                loss = lossf(pred, ytr[ix])
111                opt.zero_grad(); loss.backward(); opt.step()
112        model.eval()
113        with torch.no_grad():
114            raw = model(d["xte"].to(device))
115            pred = filter_output(raw, method)
116            metric = float(lossf(pred, d["yte"].to(device)))
117        return metric
118    except RuntimeError:
119        # Explicit CPU fallback, including shared-GPU OOM/cuDNN failures.
120        seed_all(seed)
121        d = get_dataset("dynamics", seed, n_train=400, n_test=200)
122        model = make_model("rnn_small", d["input_shape"], d["out_dim"])
123        opt = torch.optim.Adam(model.parameters(), lr=lr)
124        for _ in range(EPOCHS):
125            perm = torch.randperm(len(d["xtr"]))
126            for i in range(0, len(perm), BATCH):
127                ix = perm[i:i+BATCH]; raw = model(d["xtr"][ix])
128                loss = lossf(filter_output(raw, method), d["ytr"][ix])
129                opt.zero_grad(); loss.backward(); opt.step()
130        with torch.no_grad():
131            raw = model(d["xte"]); return float(lossf(filter_output(raw, method), d["yte"]))
132
133
134def train_fn(cfg, method):
135    return lambda seed: train_one(seed, float(cfg["lr"]), method)
136
137
138def signature(seed=0, lr=0.003):
139    """Measure behavior on a trained idea model, not a synthetic matrix."""
140    seed_all(seed)
141    d = get_dataset("dynamics", seed, n_train=400, n_test=200)
142    model = make_model("rnn_small", d["input_shape"], d["out_dim"])
143    opt = torch.optim.Adam(model.parameters(), lr=lr); lossf = nn.MSELoss()
144    for _ in range(EPOCHS):
145        p = torch.randperm(len(d["xtr"]))
146        for i in range(0, len(p), BATCH):
147            ix=p[i:i+BATCH]; raw=model(d["xtr"][ix]); loss=lossf(filter_output(raw,"idea"),d["ytr"][ix])
148            opt.zero_grad(); loss.backward(); opt.step()
149    model.eval()
150    with torch.no_grad():
151        x=d["xte"]; raw=model(x); filt=filter_output(raw,"idea")
152        # observed suppression measured on actual trained predictions
153        observed=float(torch.linalg.vector_norm(filt)/ (torch.linalg.vector_norm(raw)+1e-12))
154        t=2*GAP-1
155        f_pred=float(sum(fejer_coeffs(K)[j] * (1-t**(j+1))/(1-t) * GAP for j in range(K+1)))
156        j_pred=float(sum(jackson_coeffs(K)[j] * (1-t**(j+1))/(1-t) * GAP for j in range(len(jackson_coeffs(K)))))
157    return {"gap_hat": GAP, "K": K, "sK": GAP*K, "predicted_fejer_gain": f_pred,
158            "predicted_jackson_gain": j_pred, "observed_filtered_over_raw_norm": observed,
159            "prediction_error_abs": abs(observed-j_pred), "confirmed": abs(observed-j_pred) < .08}
160
161
162def main():
163    checks=math_checks()
164    assert checks["iterate_identity_error"] < 1e-10
165    grid=[{"lr":1e-3},{"lr":3e-3},{"lr":1e-2}]
166    base=sweep_baseline(lambda cfg: train_fn(cfg,"baseline"), grid, seeds=SWEEP_SEEDS)
167    idea_trials=[]
168    for cfg in grid:
169        r=evaluate(train_fn(cfg,"idea"), seeds=SEEDS)
170        idea_trials.append({"cfg":cfg,"result":r})
171    best=min(idea_trials,key=lambda q:q["result"]["mean"])
172    report=make_report("dynamics","rnn_small",base,best["result"],{
173        "prediction": "Jackson filtering suppresses the trained model output by p(2*gap-1) around its fixed point",
174        "signature": signature(0, float(best["cfg"]["lr"])),
175        "idea_sweep": idea_trials,
176        "math_checks": checks,
177        "protocol": {"epochs":EPOCHS,"n_train":400,"n_test":200,"gap":GAP,"K":K,
178                     "matched_architecture":True,"lr_union":grid}
179    })
180    Path("bench_report.json").write_text(json.dumps(report,indent=2))
181    print(json.dumps(report,indent=2))
182
183if __name__ == "__main__": main()