import sys, json, random import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import train_model, sweep_baseline, make_report from bench.protocol import evaluate META = {'name':'heavy_tailed_binary_radial','domain':'binary_classification','description':'Binary classification of two heavy-tailed elliptical classes, designed for learned Mahalanobis radial likelihood heads.'} def get_dataset(seed, n_train=400, n_test=400): def sample(n, s): r=np.random.RandomState(s); y=r.randint(0,2,n).astype(np.int64) mu=np.where(y[:,None]==0,[-1.15,0.0],[1.15,0.0]) # Student-t-like scale mixture, with a mild class-specific anisotropy. df=2.5 q=np.sqrt(r.chisquare(df,n)/df)[:,None] raw=r.randn(n,2)/q scale=np.where(y[:,None]==0,[1.0,.68],[1.0,.82]) x=(mu+raw*scale).astype(np.float32) return x,y xtr,ytr=sample(n_train,seed); xte,yte=sample(n_test,seed+5000) return {'xtr':xtr,'ytr':ytr,'xte':xte,'yte':yte,'task':'classification','metric':'nll'} class Encoder(nn.Module): def __init__(self): super().__init__(); self.net=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,8),nn.Tanh()) def forward(self,x): return self.net(x) class LinearSystem(nn.Module): def __init__(self): super().__init__(); self.encoder=Encoder(); self.head=nn.Linear(8,2) def forward(self,x): return self.head(self.encoder(x)) class FractionalRadialHead(nn.Module): def __init__(self,d=8,eps=1e-3): super().__init__(); self.eps=eps; self.exponents=(.25,.5,1.,1.5,2.) 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)) self.b=nn.Parameter(torch.tensor([0.,0.,-.5,0.,0.])) self.bias=nn.Parameter(torch.zeros(())) def factors(self): return torch.tril(self.off,-1)+torch.diag_embed(F.softplus(self.raw_diag)+1e-3) def covs(self): L=self.factors(); return L@L.transpose(-1,-2)+self.eps*torch.eye(L.shape[-1],device=L.device) def radii(self,z): C=torch.linalg.cholesky(self.covs()); out=[] for c in range(2): v=torch.linalg.solve_triangular(C[c],(z-self.mu[c]).T,upper=False).T; out.append((v*v).sum(1)) return torch.stack(out,1).clamp_min(1e-8) def forward(self,z): r=self.radii(z); p=torch.stack([r**a for a in self.exponents],-1); h=(p*self.b).sum(-1) ld=torch.linalg.slogdet(self.covs())[1] l=h[:,1]-h[:,0]+self.bias+.5*(ld[0]-ld[1]) return torch.stack([-l*.5,l*.5],1) class RadialSystem(nn.Module): def __init__(self): super().__init__(); self.encoder=Encoder(); self.head=FractionalRadialHead() def forward(self,x): return self.head(self.encoder(x)) def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def run_one(kind, seed, lr, epochs=28): seed_all(seed); d=get_dataset(seed) model=LinearSystem() if kind=='baseline' else RadialSystem() # The custom-track contract is numpy; bench.train_model consumes tensors. td=dict(d) for k in ('xtr','ytr','xte','yte'): td[k]=torch.from_numpy(d[k]) # train_model is the canonical path; only the readout differs. model,metric,hist=train_model(model,td,epochs=epochs,lr=lr,batch=64,weight_decay=1e-4) return float(metric), model, d def metric_fn(kind, lr, seeds, collect=False): vals=[]; models=[] for s in seeds: v,m,d=run_one(kind,s,lr) vals.append(v); models.append((m,d)) return (vals,models) if collect else vals def main(): # Cheap numerical verification: triangular solve equals explicit quadratic form. 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) v=torch.linalg.solve_triangular(C,z.T,upper=False).T; direct=(z@torch.linalg.inv(S)*z).sum(1) math_check={'max_abs_radius_error':float((v.square().sum(1)-direct).abs().max()),'finite':bool(torch.isfinite(v).all())} lrs=[0.003,0.01,0.03]; grid=[{'lr':x} for x in lrs] base=sweep_baseline(lambda c: lambda s: run_one('baseline',s,c['lr'])[0],grid,seeds=(0,1,2,3)) # Explicitly evaluate every shared lr for baseline and idea; best idea is selected on the same four tuning seeds. idea_sweep=[] for c in grid: vals=metric_fn('radial',c['lr'],(0,1,2,3)); idea_sweep.append({'cfg':c,'mean':float(np.mean(vals))}) best=min(idea_sweep,key=lambda x:x['mean'])['cfg'] idea_vals,models=metric_fn('radial',best['lr'],tuple(range(8)),True) # Signature from trained radial systems: does a nonlinear fractional basis explain their observed logits better than affine radius difference? residuals=[]; tail_slopes=[] for m,d in models: dev=next(m.parameters()).device with torch.no_grad(): 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] 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) cf=torch.linalg.lstsq(X,logits[:,None]).solution; cn=torch.linalg.lstsq(Xf,logits[:,None]).solution e0=((X@cf-logits[:,None])**2).mean().sqrt(); e1=((Xf@cn-logits[:,None])**2).mean().sqrt(); residuals.append((float(e0),float(e1))) q=torch.quantile(torch.maximum(r[:,0],r[:,1]),.95); tail_slopes.append(float(logits[torch.maximum(r[:,0],r[:,1])>=q].abs().mean())) 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]))} 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}}) with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()