Separable Ky-Fan spectral regularization / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, time, math, random
  2import numpy as np
  3import torch
  4
  5SEED = 1729
  6random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  7torch.set_num_threads(4)
  8
  9def kron_all(xs):
 10    z = xs[0]
 11    for x in xs[1:]: z = torch.kron(z, x)
 12    return z
 13
 14def numpy_kron_all(xs):
 15    z = xs[0]
 16    for x in xs[1:]: z = np.kron(z, x)
 17    return z
 18
 19def verify_majorization(trials=100, n=3, d=3):
 20    worst = -np.inf; failures = 0; max_k_gap = np.zeros(numpy_kron_all([np.eye(d)]*n).shape[0])
 21    tight_ratios=[]
 22    for t in range(trials):
 23        As=[]; Bs=[]
 24        for i in range(n):
 25            x=np.random.randn(d,d); As.append(x@x.T + .05*np.eye(d))
 26            x=np.random.randn(d,d); Bs.append(x@x.T + .05*np.eye(d))
 27        M=numpy_kron_all(As)+numpy_kron_all(Bs)
 28        lm=np.linalg.eigvalsh(M)[::-1]
 29        av=[np.linalg.eigvalsh(x)[::-1] for x in As]
 30        bv=[np.linalg.eigvalsh(x)[::-1] for x in Bs]
 31        c=numpy_kron_all(av)+numpy_kron_all(bv); c=np.sort(c)[::-1]
 32        gaps=np.cumsum(lm)-np.cumsum(c)
 33        worst=max(worst, float(gaps.max()))
 34        max_k_gap=np.maximum(max_k_gap, gaps)
 35        failures += int(np.any(gaps > 2e-9))
 36        tight_ratios.append(lm[0]/c[0])
 37    return {'trials':trials, 'failures':failures, 'worst_cumulative_gap':worst,
 38            'max_gap_by_k':max_k_gap.tolist(), 'max_eigenvalue_ratio_mean':float(np.mean(tight_ratios)),
 39            'max_eigenvalue_ratio_min':float(np.min(tight_ratios)), 'max_eigenvalue_ratio_max':float(np.max(tight_ratios))}
 40
 41class SepPSD(torch.nn.Module):
 42    def __init__(self, n=3, d=3, eps=1e-3):
 43        super().__init__(); self.n=n; self.d=d; self.eps=eps
 44        self.LA=torch.nn.ParameterList([torch.nn.Parameter(.15*torch.randn(d,d)) for _ in range(n)])
 45        self.LB=torch.nn.ParameterList([torch.nn.Parameter(.15*torch.randn(d,d)) for _ in range(n)])
 46    def factors(self):
 47        I=torch.eye(self.d, device=self.LA[0].device)
 48        A=[L@L.T+self.eps*I for L in self.LA]
 49        B=[L@L.T+self.eps*I for L in self.LB]
 50        return A,B
 51    def bound(self, k=1):
 52        A,B=self.factors()
 53        ae=[torch.linalg.eigvalsh(x).flip(0) for x in A]
 54        be=[torch.linalg.eigvalsh(x).flip(0) for x in B]
 55        c=kron_all(ae)+kron_all(be)
 56        return torch.sort(c,descending=True).values[:k].sum()/k
 57    def matrix(self):
 58        A,B=self.factors(); return kron_all(A)+kron_all(B)
 59    def forward(self,x, control=None):
 60        M=self.matrix()
 61        if control is not None:
 62            r=self.bound(1)
 63            scale=torch.clamp(torch.as_tensor(control,device=x.device)/r, max=1.0)
 64            M=M*scale
 65        return x@M.T
 66
 67def train(mode, seed=1729, steps=250, lr=.12):
 68    torch.manual_seed(seed)
 69    dev='cuda' if torch.cuda.is_available() else 'cpu'
 70    try:
 71        model=SepPSD().to(dev)
 72        g=torch.Generator(device=dev); g.manual_seed(seed+1)
 73        X=torch.randn(512,27,device=dev,generator=g)
 74        # A fixed, well-conditioned PSD target gives a nontrivial operator fitting task.
 75        q,_=torch.linalg.qr(torch.randn(27,27,device=dev,generator=g))
 76        target=q@torch.diag(torch.linspace(.15,1.0,27,device=dev))@q.T
 77        Y=X@target.T
 78        opt=torch.optim.Adam(model.parameters(),lr=lr)
 79        losses=[]; norms=[]; bounds=[]; spikes=0
 80        t0=time.perf_counter()
 81        for step in range(steps):
 82            ix=torch.arange((step*64)%448,(step*64)%448+64,device=dev)
 83            xb,yb=X[ix],Y[ix]
 84            pred=model(xb, control=2.0 if mode=='controller' else None)
 85            loss=((pred-yb)**2).mean()
 86            if mode=='penalty': loss=loss + .01*model.bound(1)
 87            if not torch.isfinite(loss): spikes += 1; break
 88            opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1000.0); opt.step()
 89            losses.append(float(loss.detach().cpu()))
 90            with torch.no_grad():
 91                M=model.matrix(); norms.append(float(torch.linalg.eigvalsh(M)[-1].cpu())); bounds.append(float(model.bound(1).cpu()))
 92        elapsed=time.perf_counter()-t0
 93        final_loss=losses[-1] if losses else float('inf')
 94        return {'mode':mode,'device':dev,'final_loss':final_loss,'best_loss':min(losses) if losses else float('inf'),
 95                'steps_completed':len(losses),'nonfinite_steps':spikes,'max_exact_norm':max(norms) if norms else None,
 96                'final_exact_norm':norms[-1] if norms else None,'final_R1':bounds[-1] if bounds else None,
 97                'norm_over_R1':norms[-1]/bounds[-1] if bounds else None,'seconds':elapsed}
 98    except Exception as e:
 99        if dev=='cuda':
100            torch.cuda.empty_cache()
101            # retry CPU by temporarily hiding CUDA is awkward; report error and caller reruns subprocess-free CPU path
102        return {'mode':mode,'error':repr(e)}
103
104def main():
105    verification=verify_majorization()
106    results=[]
107    for mode in ['baseline','penalty','controller']:
108        r=train(mode)
109        if 'error' in r and r.get('device')=='cuda':
110            # CPU fallback in-process
111            old=torch.cuda.is_available
112            torch.cuda.is_available=lambda: False
113            r=train(mode)
114            torch.cuda.is_available=old
115        results.append(r)
116    out={'seed':SEED,'verification':verification,'training':results}
117    with open('results.json','w') as f: json.dump(out,f,indent=2)
118    print(json.dumps(out,indent=2))
119if __name__=='__main__': main()