Affine-symmetry-free GMM latent prior / bench_stage2.py
Unverified
1import sys, json, math, random
2from itertools import permutations
3import numpy as np
4import torch
5from torch import nn
6import torch.nn.functional as F
7
8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
9from bench import get_dataset, evaluate, sweep_baseline, make_report
10
11SEEDS = tuple(range(8))
12SWEEP_SEEDS = (0, 1, 2, 3)
13# The union is used on both sides: baseline is evaluated at every lr/eta setting.
14CONFIGS = [
15 {"lr": 0.002, "eta": 0.0},
16 {"lr": 0.004, "eta": 0.0},
17 {"lr": 0.008, "eta": 0.0},
18]
19EPOCHS = 24
20BATCH = 64
21
22class BottleneckMLP(nn.Module):
23 def __init__(self, input_dim=10, latent_dim=8, width=32):
24 super().__init__()
25 self.encoder = nn.Sequential(nn.Linear(input_dim, width), nn.ReLU(), nn.Linear(width, latent_dim))
26 self.head = nn.Sequential(nn.Linear(latent_dim, width), nn.ReLU(), nn.Linear(width, 1))
27 self.logits = nn.Parameter(torch.randn(4) * .15)
28 self.mu = nn.Parameter(torch.randn(4, latent_dim) * .35)
29 self.logvar = nn.Parameter(torch.randn(4, latent_dim) * .05)
30 def forward(self, x):
31 return self.head(self.encoder(x))
32 def latent(self, x):
33 return self.encoder(x)
34 def cov_diag(self):
35 return torch.exp(self.logvar).clamp_min(1e-4)
36
37def sym_penalty(model, eps=.7):
38 w = torch.softmax(model.logits, 0)
39 mu = model.mu
40 var = model.cov_diag()
41 total = 0.0
42 eye = list(range(4))
43 for p in permutations(eye):
44 if list(p) == eye: continue
45 q = torch.tensor(p, device=mu.device)
46 # Signature distance: mean L2 + diagonal covariance Frobenius + weight gap.
47 dist = torch.linalg.vector_norm(mu - mu[q], dim=1)
48 dist = dist + .35 * torch.linalg.vector_norm(var - var[q], dim=1) + .8 * torch.abs(w - w[q])
49 total = total + F.softplus(eps - dist).mean()
50 return total / 23.0
51
52def signature(model):
53 with torch.no_grad():
54 w = torch.softmax(model.logits, 0).cpu().numpy()
55 mu = model.mu.cpu().numpy()
56 var = model.cov_diag().cpu().numpy()
57 vals=[]
58 for p in permutations(range(4)):
59 if list(p)==list(range(4)): continue
60 d=np.linalg.norm(mu-mu[list(p)],axis=1)+.35*np.linalg.norm(var-var[list(p)],axis=1)+.8*np.abs(w-w[list(p)])
61 vals.append(float(d.mean()))
62 return {"min_signature_distance": float(min(vals)), "mean_signature_distance": float(np.mean(vals)), "weights": w.tolist()}
63
64def train(seed, cfg, return_model=False):
65 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
66 d = get_dataset("tabular", seed, n_train=400, n_test=400)
67 # CUDA is attempted, with the same CPU fallback behavior as the bench trainer.
68 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
69 try:
70 model=BottleneckMLP(d["xtr"].shape[1]).to(device)
71 opt=torch.optim.Adam(model.parameters(), lr=cfg["lr"])
72 x,y=d["xtr"].to(device),d["ytr"].to(device)
73 for _ in range(EPOCHS):
74 model.train(); perm=torch.randperm(len(x),device=device)
75 for i in range(0,len(x),BATCH):
76 ix=perm[i:i+BATCH]; pred=model(x[ix])
77 loss=F.mse_loss(pred,y[ix])
78 if cfg.get("eta",0)>0: loss=loss+cfg["eta"]*sym_penalty(model)
79 opt.zero_grad(); loss.backward(); opt.step()
80 model.eval()
81 with torch.no_grad(): metric=F.mse_loss(model(d["xte"].to(device)),d["yte"].to(device)).item()
82 return (metric,model) if return_model else metric
83 except RuntimeError:
84 # Fresh CPU run after any CUDA failure.
85 device=torch.device("cpu"); model=BottleneckMLP(d["xtr"].shape[1])
86 opt=torch.optim.Adam(model.parameters(),lr=cfg["lr"]); x,y=d["xtr"],d["ytr"]
87 for _ in range(EPOCHS):
88 perm=torch.randperm(len(x))
89 for i in range(0,len(x),BATCH):
90 ix=perm[i:i+BATCH]; loss=F.mse_loss(model(x[ix]),y[ix])
91 if cfg.get("eta",0)>0: loss=loss+cfg["eta"]*sym_penalty(model)
92 opt.zero_grad(); loss.backward(); opt.step()
93 with torch.no_grad(): metric=F.mse_loss(model(d["xte"]),d["yte"]).item()
94 return (metric,model) if return_model else metric
95
96def main():
97 # Baseline sweep includes the complete union of learning rates tried by the idea.
98 base=sweep_baseline(lambda c: (lambda seed: train(seed,c)), CONFIGS, seeds=SWEEP_SEEDS)
99 best_lr=base["best_cfg"]["lr"]
100 idea_cfgs=[{"lr":best_lr,"eta":e} for e in (.02,.08,.20)]
101 idea_results=[]
102 for c in idea_cfgs:
103 r=evaluate(lambda seed,c=c: train(seed,c), seeds=SEEDS)
104 idea_results.append((c,r))
105 best_cfg,best=min(idea_results,key=lambda z:z[1]["mean"]) # min mean
106 # Trained behavior signature, not an analytical toy identity.
107 sig_rows=[]
108 for c,_ in idea_results:
109 distances=[]
110 for s in SEEDS:
111 _,m=train(s,c,True); distances.append(signature(m)["min_signature_distance"])
112 sig_rows.append({"cfg":c,"observed_min_dist_mean":float(np.mean(distances)),"observed_min_dist_per_seed":distances})
113 # Stage-1 quantitative prediction: increasing eta should increase separation.
114 observed=[x["observed_min_dist_mean"] for x in sig_rows]
115 confirmed=bool(all(observed[i+1] > observed[i] for i in range(len(observed)-1)))
116 report=make_report("tabular","mlp_bottleneck",base,best,extra={
117 "prediction":"larger eta produces larger trained GMM signature separation",
118 "eta_rows":sig_rows,"predicted_order":"eta .02 < .08 < .20",
119 "confirmed":confirmed
120 })
121 report["idea_sweep"]= [{"cfg":c,"mean":r["mean"],"std":r["std"],"per_seed":r["per_seed"]} for c,r in idea_results]
122 report["protocol_note"]="Tabular is the built-in regularization/optimizer track; both systems use the identical bottleneck MLP and differ only by the signature loss."
123 with open("bench_report.json","w") as f: json.dump(report,f,indent=2)
124 print(json.dumps(report,indent=2))
125
126if __name__=="__main__": main()