Fractional Mahalanobis radial head / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5import torch.nn.functional as F
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import train_model, sweep_baseline, make_report
  8from bench.protocol import evaluate
  9
 10META = {'name':'heavy_tailed_binary_radial','domain':'binary_classification','description':'Binary classification of two heavy-tailed elliptical classes, designed for learned Mahalanobis radial likelihood heads.'}
 11
 12def get_dataset(seed, n_train=400, n_test=400):
 13    def sample(n, s):
 14        r=np.random.RandomState(s); y=r.randint(0,2,n).astype(np.int64)
 15        mu=np.where(y[:,None]==0,[-1.15,0.0],[1.15,0.0])
 16        # Student-t-like scale mixture, with a mild class-specific anisotropy.
 17        df=2.5
 18        q=np.sqrt(r.chisquare(df,n)/df)[:,None]
 19        raw=r.randn(n,2)/q
 20        scale=np.where(y[:,None]==0,[1.0,.68],[1.0,.82])
 21        x=(mu+raw*scale).astype(np.float32)
 22        return x,y
 23    xtr,ytr=sample(n_train,seed); xte,yte=sample(n_test,seed+5000)
 24    return {'xtr':xtr,'ytr':ytr,'xte':xte,'yte':yte,'task':'classification','metric':'nll'}
 25
 26class Encoder(nn.Module):
 27    def __init__(self):
 28        super().__init__(); self.net=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,8),nn.Tanh())
 29    def forward(self,x): return self.net(x)
 30
 31class LinearSystem(nn.Module):
 32    def __init__(self):
 33        super().__init__(); self.encoder=Encoder(); self.head=nn.Linear(8,2)
 34    def forward(self,x): return self.head(self.encoder(x))
 35
 36class FractionalRadialHead(nn.Module):
 37    def __init__(self,d=8,eps=1e-3):
 38        super().__init__(); self.eps=eps; self.exponents=(.25,.5,1.,1.5,2.)
 39        self.mu=nn.Parameter(torch.zeros(2,d)); self.raw_diag=nn.Parameter(torch.full((2,d),-0.3)); self.off=nn.Parameter(torch.zeros(2,d,d))
 40        self.b=nn.Parameter(torch.tensor([0.,0.,-.5,0.,0.]))
 41        self.bias=nn.Parameter(torch.zeros(()))
 42    def factors(self):
 43        return torch.tril(self.off,-1)+torch.diag_embed(F.softplus(self.raw_diag)+1e-3)
 44    def covs(self):
 45        L=self.factors(); return L@L.transpose(-1,-2)+self.eps*torch.eye(L.shape[-1],device=L.device)
 46    def radii(self,z):
 47        C=torch.linalg.cholesky(self.covs()); out=[]
 48        for c in range(2):
 49            v=torch.linalg.solve_triangular(C[c],(z-self.mu[c]).T,upper=False).T; out.append((v*v).sum(1))
 50        return torch.stack(out,1).clamp_min(1e-8)
 51    def forward(self,z):
 52        r=self.radii(z); p=torch.stack([r**a for a in self.exponents],-1); h=(p*self.b).sum(-1)
 53        ld=torch.linalg.slogdet(self.covs())[1]
 54        l=h[:,1]-h[:,0]+self.bias+.5*(ld[0]-ld[1])
 55        return torch.stack([-l*.5,l*.5],1)
 56
 57class RadialSystem(nn.Module):
 58    def __init__(self):
 59        super().__init__(); self.encoder=Encoder(); self.head=FractionalRadialHead()
 60    def forward(self,x): return self.head(self.encoder(x))
 61
 62def seed_all(s):
 63    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 64    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 65
 66def run_one(kind, seed, lr, epochs=28):
 67    seed_all(seed); d=get_dataset(seed)
 68    model=LinearSystem() if kind=='baseline' else RadialSystem()
 69    # The custom-track contract is numpy; bench.train_model consumes tensors.
 70    td=dict(d)
 71    for k in ('xtr','ytr','xte','yte'):
 72        td[k]=torch.from_numpy(d[k])
 73    # train_model is the canonical path; only the readout differs.
 74    model,metric,hist=train_model(model,td,epochs=epochs,lr=lr,batch=64,weight_decay=1e-4)
 75    return float(metric), model, d
 76
 77def metric_fn(kind, lr, seeds, collect=False):
 78    vals=[]; models=[]
 79    for s in seeds:
 80        v,m,d=run_one(kind,s,lr)
 81        vals.append(v); models.append((m,d))
 82    return (vals,models) if collect else vals
 83
 84def main():
 85    # Cheap numerical verification: triangular solve equals explicit quadratic form.
 86    torch.manual_seed(7); A=torch.randn(4,4,dtype=torch.double); S=A@A.T+.2*torch.eye(4,dtype=torch.double); C=torch.linalg.cholesky(S); z=torch.randn(6,4,dtype=torch.double)
 87    v=torch.linalg.solve_triangular(C,z.T,upper=False).T; direct=(z@torch.linalg.inv(S)*z).sum(1)
 88    math_check={'max_abs_radius_error':float((v.square().sum(1)-direct).abs().max()),'finite':bool(torch.isfinite(v).all())}
 89    lrs=[0.003,0.01,0.03]; grid=[{'lr':x} for x in lrs]
 90    base=sweep_baseline(lambda c: lambda s: run_one('baseline',s,c['lr'])[0],grid,seeds=(0,1,2,3))
 91    # Explicitly evaluate every shared lr for baseline and idea; best idea is selected on the same four tuning seeds.
 92    idea_sweep=[]
 93    for c in grid:
 94        vals=metric_fn('radial',c['lr'],(0,1,2,3)); idea_sweep.append({'cfg':c,'mean':float(np.mean(vals))})
 95    best=min(idea_sweep,key=lambda x:x['mean'])['cfg']
 96    idea_vals,models=metric_fn('radial',best['lr'],tuple(range(8)),True)
 97    # Signature from trained radial systems: does a nonlinear fractional basis explain their observed logits better than affine radius difference?
 98    residuals=[]; tail_slopes=[]
 99    for m,d in models:
100        dev=next(m.parameters()).device
101        with torch.no_grad():
102            z=m.encoder(torch.tensor(d['xte'],device=dev)); logits=m.head(z)[:,1]-m.head(z)[:,0]; r=m.head.radii(z); x=r[:,1]-r[:,0]
103            X=torch.stack([torch.ones_like(x),x],1); Xf=torch.stack([torch.ones_like(x)]+[r[:,1]**a-r[:,0]**a for a in m.head.exponents],1)
104            cf=torch.linalg.lstsq(X,logits[:,None]).solution; cn=torch.linalg.lstsq(Xf,logits[:,None]).solution
105            e0=((X@cf-logits[:,None])**2).mean().sqrt(); e1=((Xf@cn-logits[:,None])**2).mean().sqrt(); residuals.append((float(e0),float(e1)))
106            q=torch.quantile(torch.maximum(r[:,0],r[:,1]),.95); tail_slopes.append(float(logits[torch.maximum(r[:,0],r[:,1])>=q].abs().mean()))
107    sig={'prediction':'fractional radial readout produces measurable nonlinear curvature beyond affine radius difference on trained heavy-tailed systems','affine_logit_rmse_mean':float(np.mean([x[0] for x in residuals])),'fractional_basis_logit_rmse_mean':float(np.mean([x[1] for x in residuals])),'tail_abs_logit_mean':float(np.mean(tail_slopes)),'confirmed':bool(np.mean([x[1] for x in residuals]) < .95*np.mean([x[0] for x in residuals]))}
108    rep=make_report('custom:heavy_tailed_binary_radial','mlp_tiny',base,{'mean':float(np.mean(idea_vals)),'std':float(np.std(idea_vals)),'per_seed':idea_vals,'n':8},extra={'mechanism_signature':sig,'custom_track':{'name':META['name'],'file':'bench_experiment.py','domain':META['domain']},'math_check':math_check,'idea_sweep':idea_sweep,'selection':{'best_idea_cfg':best,'shared_lr_grid':lrs}})
109    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
110    print(json.dumps(rep,indent=2))
111if __name__=='__main__': main()