Fractional Mahalanobis radial head / bench_experiment.py
Mechanism confirmed, baseline not beaten
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()