Separable Ky-Fan spectral regularization / bench_run.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1from __future__ import annotations
  2import json, random, sys
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10EPOCHS = 15
 11BATCH = 128
 12
 13def kron_np(xs):
 14    z = xs[0]
 15    for x in xs[1:]: z = np.kron(z, x)
 16    return z
 17
 18def verify_majorization(trials=40, n=3, d=3):
 19    rng = np.random.default_rng(123)
 20    failures, worst, ratios = 0, -float('inf'), []
 21    for _ in range(trials):
 22        aa, bb = [], []
 23        for _ in range(n):
 24            q = rng.normal(size=(d,d)); aa.append(q @ q.T + .05*np.eye(d))
 25            q = rng.normal(size=(d,d)); bb.append(q @ q.T + .05*np.eye(d))
 26        lm = np.linalg.eigvalsh(kron_np(aa)+kron_np(bb))[::-1]
 27        av = [np.linalg.eigvalsh(x)[::-1] for x in aa]
 28        bv = [np.linalg.eigvalsh(x)[::-1] for x in bb]
 29        c = np.sort(kron_np(av)+kron_np(bv))[::-1]
 30        gap = float(np.max(np.cumsum(lm)-np.cumsum(c)))
 31        worst = max(worst, gap); failures += int(gap > 2e-8)
 32        ratios.append(float(lm[0]/c[0]))
 33    return {'trials':trials, 'failures':failures,
 34            'worst_cumulative_gap':worst,
 35            'max_eigenvalue_ratio_mean':float(np.mean(ratios))}
 36
 37def kron_t(xs):
 38    z = xs[0]
 39    for x in xs[1:]: z = torch.kron(z, x)
 40    return z
 41
 42class PSDKron64(nn.Module):
 43    """Shared MLP with its first 64x64 linear map represented as two PSD terms."""
 44    def __init__(self, input_dim, out_dim, eps=1e-3):
 45        super().__init__(); self.eps = eps
 46        self.inp = nn.Linear(input_dim, 64)
 47        self.LA = nn.ParameterList([nn.Parameter(.08*torch.randn(8,8)) for _ in range(2)])
 48        self.LB = nn.ParameterList([nn.Parameter(.08*torch.randn(8,8)) for _ in range(2)])
 49        self.bias = nn.Parameter(torch.zeros(64))
 50        self.out = nn.Linear(64, out_dim)
 51    def factors(self):
 52        dev=self.LA[0].device; I=torch.eye(8,device=dev)
 53        return ([x@x.T+self.eps*I for x in self.LA], [x@x.T+self.eps*I for x in self.LB])
 54    def bound(self,k=1):
 55        A,B=self.factors()
 56        av=[torch.linalg.eigvalsh(x).flip(0) for x in A]
 57        bv=[torch.linalg.eigvalsh(x).flip(0) for x in B]
 58        c=torch.sort(kron_t(av)+kron_t(bv),descending=True).values
 59        return c[:k].mean()
 60    def exact_norm(self):
 61        A,B=self.factors(); return torch.linalg.eigvalsh(kron_t(A)+kron_t(B))[-1]
 62    def forward(self,x):
 63        h=torch.relu(self.inp(x))
 64        A,B=self.factors(); M=kron_t(A)+kron_t(B)
 65        h=torch.relu(h@M.T+self.bias)
 66        return self.out(h)
 67
 68def seed_all(seed):
 69    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 70
 71def baseline_one(cfg, seed, collect=False):
 72    seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=200)
 73    net=make_model('mlp_tiny',ds['input_shape'],ds['out_dim'])
 74    net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['wd'],log=lambda _:None)
 75    return float(metric)
 76
 77def idea_one(cfg, seed, collect=False):
 78    seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=200)
 79    dev='cuda' if torch.cuda.is_available() else 'cpu'
 80    try:
 81        net=PSDKron64(int(np.prod(ds['input_shape'])),ds['out_dim']).to(dev)
 82        xtr,ytr=ds['xtr'].to(dev),ds['ytr'].to(dev); opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['wd'])
 83        vals=[]; bounds=[]; norms=[]
 84        for ep in range(EPOCHS):
 85            net.train(); perm=torch.randperm(len(xtr),device=dev)
 86            for i in range(0,len(xtr),BATCH):
 87                ix=perm[i:i+BATCH]; pred=net(xtr[ix]); loss=((pred-ytr[ix])**2).mean()+cfg['lam']*net.bound(cfg['k'])
 88                opt.zero_grad(); loss.backward(); opt.step(); vals.append(float(loss.detach().cpu()))
 89            with torch.no_grad(): bounds.append(float(net.bound(cfg['k']).cpu())); norms.append(float(net.exact_norm().cpu()))
 90        net.eval()
 91        with torch.no_grad(): metric=float(((net(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean().cpu())
 92        if collect:
 93            b1=[]
 94            with torch.no_grad():
 95                for _ in norms: b1.append(float(net.bound(1).cpu()))
 96            return metric, {'bound_R4':float(np.mean(bounds)),'bound_R1':float(np.mean(b1)),'exact_norm':float(np.mean(norms)),'ratio_R1':float(np.mean(np.array(norms)/np.array(b1)))}
 97        return metric
 98    except Exception:
 99        if dev=='cuda':
100            torch.cuda.empty_cache(); torch.set_default_device('cpu'); return idea_one(cfg,seed,collect)
101        raise
102
103def main():
104    verification=verify_majorization()
105    # Union parity: every idea lr/wd is also evaluated by baseline.
106    grid=[{'lr':lr,'wd':wd,'lam':lam,'k':k} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4) for lam,k in ((0.0,1),)]
107    basegrid=[{'lr':x['lr'],'wd':x['wd']} for x in grid]
108    base=sweep_baseline(lambda c: (lambda s: baseline_one(c,s)),basegrid,seeds=tuple(range(4)))
109    # Baseline sweep includes all idea learning rates; final best is evaluated on 8 seeds.
110    idea_grid=[{'lr':lr,'wd':wd,'lam':lam,'k':k} for lr in (0.001,0.003,0.006) for wd in (0.0,1e-4) for lam,k in ((0.001,1),(0.003,1),(0.001,4))]
111    tried=[]
112    for c in idea_grid:
113        r=evaluate(lambda s: idea_one(c,s),seeds=tuple(range(4))); tried.append({'cfg':c,'mean':r['mean']})
114    best=min(tried,key=lambda x:x['mean'])['cfg']
115    idea=evaluate(lambda s: idea_one(best,s),seeds=SEEDS)
116    # Trained-model signature, independently measured on two paired trained models.
117    sig=[]
118    for s in (0,1):
119        _,z=idea_one(best,s,True); sig.append(z)
120    signature={'quantity':'exact operator norm / separable R1 bound on trained models','predicted':'ratio <= 1 by the Ky-Fan upper bound','observed_mean_ratio':float(np.mean([z['ratio_R1'] for z in sig])),'observed':sig,'confirmed':all(z['ratio_R1']<=1.0001 for z in sig)}
121    base['idea_union_sweep']=tried
122    rep=make_report('tabular','mlp_tiny',base,idea,{'mechanism_signature':signature,'verification':verification,'idea_best_cfg':best})
123    rep['stage2_note']='Baseline and idea use same tabular data, epochs, batch size, Adam, learning-rate/weight-decay union; idea changes only the first hidden operator and adds Ky-Fan penalty.'
124    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
125    print(json.dumps(rep,indent=2))
126if __name__=='__main__': main()