Interlevel Betti Token Transformer / bench_graph.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import sys, json
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  8from bench import train_model, sweep_baseline, evaluate, make_report
  9from graph_topology_track import get_dataset
 10
 11SEEDS = tuple(range(8))
 12TUNE_SEEDS = (0, 1, 2, 3)
 13LRS = [0.003, 0.01, 0.03]
 14GRID = np.linspace(0.0, 1.0, 9, dtype=np.float32)
 15M, STRIDE, N = 4, 2, 12
 16
 17
 18def gf2_rank(a):
 19    a = (np.asarray(a, dtype=np.uint8) & 1).copy()
 20    r = 0
 21    for c in range(a.shape[1]):
 22        piv = np.flatnonzero(a[r:, c])
 23        if len(piv) == 0: continue
 24        p = r + int(piv[0]); a[[r, p]] = a[[p, r]]
 25        for q in range(a.shape[0]):
 26            if q != r and a[q, c]: a[q] ^= a[r]
 27        r += 1
 28        if r == a.shape[0]: break
 29    return r
 30
 31
 32def betti(nv, edges):
 33    B = np.zeros((nv, len(edges)), dtype=np.uint8)
 34    for j, (u, v) in enumerate(edges): B[u, j] = B[v, j] = 1
 35    rank = gf2_rank(B)
 36    return int(nv - rank), int(len(edges) - rank)
 37
 38
 39def tokens(x, topo):
 40    out = []
 41    for row in x:
 42        A, h = row[:N*N].reshape(N, N), row[N*N:]
 43        edges = [(i, j) for i in range(N) for j in range(i + 1, N) if A[i, j] > .5]
 44        seq = []
 45        for st in range(0, len(GRID) - M, STRIDE):
 46            lo, hi = float(GRID[st]), float(GRID[st + M])
 47            ee = [(u, v) for u, v in edges if lo <= max(h[u], h[v]) <= hi]
 48            keep = {i for i, z in enumerate(h) if lo <= z <= hi}
 49            for u, v in ee: keep.update((u, v))
 50            vs = sorted(keep); rem = {v: i for i, v in enumerate(vs)}
 51            redges = [(rem[u], rem[v]) for u, v in ee]
 52            b0, b1 = betti(len(vs), redges)
 53            if topo:
 54                seq.append([b0, b1, np.log1p(len(vs)), np.log1p(len(redges))])
 55            else:
 56                deg = np.zeros(len(vs), dtype=np.float32)
 57                for u, v in redges: deg[u] += 1; deg[v] += 1
 58                seq.append([np.mean(deg) if len(deg) else 0., np.std(deg) if len(deg) else 0.,
 59                            np.log1p(len(vs)), np.log1p(len(redges))])
 60        out.append(seq)
 61    return np.asarray(out, dtype=np.float32)
 62
 63
 64class TokenTransformer(nn.Module):
 65    def __init__(self, n_tokens, width=32):
 66        super().__init__()
 67        self.proj = nn.Linear(4, width)
 68        self.pos = nn.Parameter(torch.zeros(1, n_tokens, width))
 69        nn.init.normal_(self.pos, std=.02)
 70        layer = nn.TransformerEncoderLayer(width, 4, 2 * width, batch_first=True, dropout=0.)
 71        self.enc = nn.TransformerEncoder(layer, 1)
 72        self.head = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, 2))
 73    def forward(self, x):
 74        z = self.proj(x) + self.pos[:, :x.shape[1]]
 75        return self.head(self.enc(z).mean(1))
 76
 77
 78def make_ds(raw, topo):
 79    return {"xtr": torch.tensor(tokens(raw["xtr"], topo)), "ytr": torch.tensor(raw["ytr"], dtype=torch.long),
 80            "xte": torch.tensor(tokens(raw["xte"], topo)), "yte": torch.tensor(raw["yte"], dtype=torch.long),
 81            "task": "classification", "metric": "err"}
 82
 83
 84def run_one(raw, topo, lr, seed):
 85    torch.manual_seed(seed + (10000 if topo else 0)); np.random.seed(seed)
 86    ds = make_ds(raw, topo)
 87    net, metric, _ = train_model(TokenTransformer(ds["xtr"].shape[1]), ds, epochs=18, lr=float(lr), batch=128, log=lambda *_: None)
 88    return float(metric), net, ds
 89
 90
 91def factory(topo, lr):
 92    def fn(seed):
 93        raw = get_dataset(int(seed), 400, 200)
 94        return run_one(raw, topo, lr, int(seed))[0]
 95    return fn
 96
 97
 98def main():
 99    grid = [{"lr": lr} for lr in LRS]
100    baseline = sweep_baseline(lambda cfg: factory(False, cfg["lr"]), grid, seeds=TUNE_SEEDS)
101    idea_trials = [{"cfg": cfg, "mean": evaluate(factory(True, cfg["lr"]), seeds=TUNE_SEEDS)["mean"]} for cfg in grid]
102    idea_cfg = min(idea_trials, key=lambda z: z["mean"])["cfg"]
103    idea_lr = float(idea_cfg["lr"])
104    idea_full = evaluate(factory(True, idea_lr), seeds=SEEDS)
105    pred_acc, b1_corr = [], []
106    for seed in SEEDS:
107        raw = get_dataset(seed, 400, 200)
108        _, net, ds = run_one(raw, True, idea_lr, seed)
109        dev = next(net.parameters()).device
110        net.eval()
111        with torch.no_grad(): pred = net(ds["xte"].to(dev)).argmax(1).cpu().numpy()
112        observed_b1 = []
113        for row in raw["xte"]:
114            A = row[:N*N].reshape(N, N)
115            edges = [(i, j) for i in range(N) for j in range(i+1, N) if A[i,j] > .5]
116            observed_b1.append(betti(N, edges)[1])
117        observed_b1 = np.asarray(observed_b1)
118        pred_acc.append(float(np.mean(pred == raw["yte"])))
119        b1_corr.append(float(np.corrcoef(pred, observed_b1)[0, 1]))
120    signature = {"predicted": {"class0_betti1": 1.0, "class1_betti1": 2.0},
121                 "observed": {"idea_accuracy_mean": float(np.mean(pred_acc)),
122                              "prediction_observed_betti1_correlation_mean": float(np.mean(b1_corr))},
123                 "confirmed": bool(np.mean(pred_acc) >= .75 and np.mean(b1_corr) > .5),
124                 "note": "Measured from eight independently trained idea models on held-out graphs."}
125    base_block = {"best_cfg": baseline["best_cfg"], "sweep": baseline["sweep"], "full": baseline["full"]}
126    idea = {"best_cfg": idea_cfg, "sweep": idea_trials, "mean": idea_full["mean"],
127            "std": idea_full["std"], "per_seed": idea_full["per_seed"], "n": idea_full["n"]}
128    report = make_report("cycle_topology_graph", "token_transformer", base_block, idea, signature)
129    report["custom_track"] = {"name": "cycle_topology_graph", "file": "graph_topology_track.py", "domain": "graph_topology"}
130    report["protocol_note"] = "Eight paired seeds; baseline and idea share the same Transformer architecture and lr union; only token construction differs."
131    print(json.dumps(report, indent=2))
132
133if __name__ == "__main__": main()