GQL Safe Residual Layer / stage2_bench.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import evaluate, sweep_baseline, make_report
7from gql_track import get_dataset
8
9SEEDS = tuple(range(8))
10GRID = [{'lr': 1e-3, 'epochs': 12}, {'lr': 3e-3, 'epochs': 12}, {'lr': 6e-3, 'epochs': 12}]
11EPS = 1e-4
12
13class MLP(nn.Module):
14 def __init__(self):
15 super().__init__()
16 self.net = nn.Sequential(nn.Linear(8, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 4))
17 def forward(self, x):
18 return self.net(x)
19
20def gql_limit(u0, cand, eps=EPS, iters=60):
21 delta = cand - u0
22 def margin(t):
23 v = u0 + t.unsqueeze(-1) * delta
24 return v[:, 3] - torch.sqrt(torch.sum(v[:, :3] * v[:, :3], dim=1) + 1e-12)
25 one = torch.ones(u0.shape[0], device=u0.device)
26 bad = (cand[:, 0] < eps) | (margin(one) < eps)
27 lo = torch.zeros_like(one); hi = one
28 for _ in range(iters):
29 mid = (lo + hi) * .5
30 ok = margin(mid) >= eps
31 lo = torch.where(ok, mid, lo); hi = torch.where(ok, hi, mid)
32 theta = torch.where(bad, lo, one)
33 db = cand[:, 0] < eps
34 td = torch.clamp((u0[:, 0] - eps) / torch.clamp(u0[:, 0] - cand[:, 0], min=1e-12), 0., 1.)
35 theta = torch.where(db, torch.minimum(theta, td), theta)
36 return u0 + theta.unsqueeze(-1) * delta, theta
37
38def train(seed, cfg, limited):
39 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
40 d = get_dataset(seed, 400, 100)
41 try:
42 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
43 model = MLP().to(device)
44 xtr, ytr = torch.as_tensor(d['xtr'], device=device), torch.as_tensor(d['ytr'], device=device)
45 xte, yte = torch.as_tensor(d['xte'], device=device), torch.as_tensor(d['yte'], device=device)
46 opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'])
47 model.train()
48 for _ in range(cfg['epochs']):
49 p = model(xtr)
50 if limited:
51 p, _ = gql_limit(xtr[:, :4], p)
52 loss = ((p-ytr)**2).mean()
53 opt.zero_grad(); loss.backward(); opt.step()
54 model.eval()
55 with torch.no_grad():
56 raw = model(xte)
57 pred, theta = gql_limit(xte[:, :4], raw) if limited else (raw, torch.ones(raw.shape[0], device=device))
58 mse = ((pred-yte)**2).mean().item()
59 margin = pred[:,3]-torch.sqrt(torch.sum(pred[:,:3]**2,dim=1)+1e-12)
60 invalid = ((pred[:,0] < EPS) | (margin < EPS)).float().mean().item()
61 limited_fraction = (theta < .999999).float().mean().item()
62 return mse, {'invalid_rate':invalid,'limited_fraction':limited_fraction,'mean_theta':theta.mean().item()}
63 except Exception:
64 device=torch.device('cpu'); torch.manual_seed(seed)
65 model=MLP(); xtr=torch.as_tensor(d['xtr']); ytr=torch.as_tensor(d['ytr'])
66 opt=torch.optim.Adam(model.parameters(),lr=cfg['lr'])
67 for _ in range(cfg['epochs']):
68 p=model(xtr); p,_=gql_limit(xtr[:,:4],p) if limited else (p,None)
69 loss=((p-ytr)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
70 with torch.no_grad():
71 raw=model(torch.as_tensor(d['xte'])); pred,th=gql_limit(torch.as_tensor(d['xte'][:,:4]),raw) if limited else (raw,torch.ones(100))
72 margin=pred[:,3]-torch.sqrt(torch.sum(pred[:,:3]**2,1)+1e-12)
73 return ((pred-torch.as_tensor(d['yte']))**2).mean().item(), {'invalid_rate':float(((pred[:,0]<EPS)|(margin<EPS)).float().mean()),'limited_fraction':float((th<.999999).float().mean()),'mean_theta':float(th.mean())}
74
75def fn(cfg, limited):
76 return lambda seed: train(seed,cfg,limited)[0]
77
78def mechanism_signature():
79 rows=[]
80 for s in SEEDS:
81 cfg=GRID[1]; d=get_dataset(s,400,100); random.seed(s); np.random.seed(s); torch.manual_seed(s)
82 model=MLP(); opt=torch.optim.Adam(model.parameters(),lr=cfg['lr']); x=torch.as_tensor(d['xtr']); y=torch.as_tensor(d['ytr'])
83 for _ in range(cfg['epochs']):
84 p,_=gql_limit(x[:,:4],model(x)); loss=((p-y)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
85 with torch.no_grad():
86 raw=model(torch.as_tensor(d['xte'])); safe,th=gql_limit(torch.as_tensor(d['xte'][:,:4]),raw)
87 margin=safe[:,3]-torch.sqrt(torch.sum(safe[:,:3]**2,1)+1e-12)
88 rows.append({'seed':s,'invalid_raw':float(((raw[:,0]<EPS)|((raw[:,3]-torch.sqrt(torch.sum(raw[:,:3]**2,1)+1e-12))<EPS)).float().mean()),'invalid_limited':float(((safe[:,0]<EPS)|(margin<EPS)).float().mean()),'mean_theta':float(th.mean())})
89 pred=0.0; observed=float(np.mean([r['invalid_limited'] for r in rows])); raw=float(np.mean([r['invalid_raw'] for r in rows]))
90 return {'claim':'limiting a valid baseline candidate along the complete residual enforces admissibility','predicted_invalid_rate_after_limit':pred,'observed_invalid_rate_after_limit':observed,'observed_raw_invalid_rate':raw,'trained_model_measurements':rows,'confirmed':bool(observed <= 1e-5 and raw > observed)}
91
92def main():
93 base=sweep_baseline(lambda c: fn(c,False), GRID, seeds=(0,1,2,3))
94 idea_sweep=[{'cfg':c,'mean':evaluate(fn(c,True),seeds=(0,1,2,3))['mean']} for c in GRID]
95 best=min(idea_sweep,key=lambda z:z['mean'])['cfg']
96 idea=evaluate(fn(best,True),seeds=SEEDS)
97 rep=make_report('custom_relativistic_conservative_regression','mlp_tiny',base,idea,{'mechanism_signature':mechanism_signature(),'custom_track':{'name':'relativistic_conservative_regression','file':'gql_track.py','domain':'conservative_state_admissibility'},'sweep_parity':{'union_grid':GRID,'idea_sweep':idea_sweep}})
98 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
99 print(json.dumps(rep,indent=2))
100
101if __name__=='__main__': main()