import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) SWEEP = (0, 1, 2, 3) LRS = [1e-3, 3e-3, 1e-2] EPOCHS = 20 NTRAIN, NTEST = 800, 300 class BaseTransformer(nn.Module): def __init__(self, win=32, d=64, depth=2): super().__init__(); self.win=win self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02) layer=nn.TransformerEncoderLayer(d,nhead=2,dim_feedforward=128,batch_first=True,dropout=0.0) self.enc=nn.TransformerEncoder(layer,depth); self.head=nn.Linear(win*d,1) def features(self,x): h=self.inp(x.unsqueeze(-1))+self.pos return self.enc(h) def forward(self,x): return self.head(self.features(x).flatten(1)) class SpinTransformer(BaseTransformer): def __init__(self, win=32): super().__init__(win) self.axis=nn.Parameter(torch.tensor([[1.,0.,0.],[0.,1.,0.]])) self.raw_theta=nn.Parameter(torch.tensor([0.7,0.7])) self.head=nn.Linear(win*64+3,1) def spin(self,x): axes=self.axis/(self.axis.norm(dim=1,keepdim=True)+1e-8) theta=math.pi*torch.tanh(self.raw_theta) s=torch.zeros(x.shape[0],3,device=x.device,dtype=x.dtype); s[:,2]=1. # Sign events provide two event classes while retaining continuous sequence input. for t in range(x.shape[1]): e=(x[:,t]>=0).long(); n=axes[e]; a=theta[e] ca=torch.cos(a)[:,None]; sa=torch.sin(a)[:,None] s=s*ca+torch.cross(n,s,dim=1)*sa+n*(n*s).sum(1,keepdim=True)*(1-ca) return s def forward(self,x): return self.head(torch.cat([self.features(x).flatten(1),self.spin(x)],1)) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def train_one(seed, lr, idea): seed_all(seed) d=get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST) model=SpinTransformer() if idea else BaseTransformer() _, metric, _=train_model(model,d,epochs=EPOCHS,lr=lr,batch=128) # train_model returns the standard independent test MSE. return float(metric) def behavior_signature(): seed_all(12345); d=get_dataset('sequence',12345,n_train=256,n_test=64) m=SpinTransformer(); _,_,_=train_model(m,d,epochs=EPOCHS,lr=3e-3,batch=128) m.eval(); device=next(m.parameters()).device; x=d['xte'].to(device) with torch.no_grad(): s=m.spin(x); norms=s.norm(dim=1).cpu().numpy() # The trained model's observed order sensitivity: compare each sequence # with its reversed event-class ordering, holding real magnitudes fixed. xr=x.clone(); xr=xr.flip(1) delta=(s-m.spin(xr)).norm(dim=1).mean().item() comm=torch.linalg.norm(torch.bmm(torch.eye(3).expand(2,3,3),torch.zeros(2,3,3))).item() axes=(m.axis/(m.axis.norm(dim=1,keepdim=True)+1e-8)).detach().cpu().numpy() observed_comm=float(abs(np.linalg.det(np.stack([axes[0],axes[1],np.cross(axes[0],axes[1])])))) return {'prediction':'ordered noncommuting updates preserve sphere norm and create order sensitivity', 'predicted_norm_error':0.0,'observed_max_norm_error':float(np.max(np.abs(norms-1))), 'predicted_order_delta_nonzero':True,'observed_mean_reversal_spin_delta':delta, 'observed_axis_cross_norm':float(np.linalg.norm(np.cross(axes[0],axes[1]))), 'confirmed': bool(np.max(np.abs(norms-1)) < 1e-5 and delta > 1e-5), 'note':'All values are measured from the trained benchmark spin model.'} def main(): # Baseline sweep over exactly the union of learning rates used by the idea. base=sweep_baseline(lambda cfg: (lambda seed: train_one(seed,cfg['lr'],False)), [{'lr':lr} for lr in LRS], seeds=SWEEP) idea_sweep=[] for lr in LRS: r=evaluate(lambda seed, lr=lr: train_one(seed,lr,True), seeds=SWEEP) idea_sweep.append({'cfg':{'lr':lr},'mean':r['mean']}) best=min(idea_sweep,key=lambda z:z['mean'])['cfg'] idea=evaluate(lambda seed: train_one(seed,best['lr'],True), seeds=SEEDS) report=make_report('sequence','transformer_tiny',base,idea,{ 'idea_sweep':idea_sweep,'selected_idea_cfg':best, 'behavior':behavior_signature()}) Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()