GQL Safe Residual Layer / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  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()