Critical-Tail Multiscale Mixer / run_experiment.py
Beats tuned baseline
1import math, json, random, time
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7SEED=417
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9try:
10 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
11 if device.type=='cuda': torch.cuda.empty_cache()
12except Exception:
13 device=torch.device('cpu')
14
15def kval(r, r0=2.0):
16 return 1.0/((r+r0)*math.log(r+r0)**2)
17
18def math_check():
19 rs=np.arange(1,300001,dtype=np.float64)
20 kv=1/((rs+2)*np.log(rs+2)**2)
21 tail=[]
22 for R in [8,16,32,64,128,256]:
23 s=float(kv[R:].sum()); q=1/math.log(R+2)
24 tail.append([R,s,q,s/q])
25 bands=[]
26 for k in range(7):
27 lo,hi=2**k,2**(k+1)
28 exact=float(kv[lo-1:hi-1].sum())
29 formula=1/((k+1)*(k+2)*math.log(2))
30 bands.append([k,exact,formula,exact/formula])
31 return {'tail_R_discrete_integral_ratio':tail,'bands_k_exact_formula_ratio':bands}
32
33def make_weights(L):
34 w=torch.zeros(L,device=device)
35 for r in range(1,L): w[r]=kval(r)
36 return w/(2*w.sum())
37
38def long_target(x,w):
39 # x [B,L,1], symmetric zero-padded convolution
40 B,L,D=x.shape; y=torch.zeros_like(x)
41 for r in range(1,L):
42 y[:,r:]+=w[r]*x[:,:L-r]; y[:,:L-r]+=w[r]*x[:,r:]
43 return y
44
45def dyadic_mix(x):
46 B,L,D=x.shape; out=torch.zeros_like(x); pref=torch.cat([torch.zeros(B,1,D,device=x.device),x.cumsum(1)],1)
47 K=int(math.log2(L))
48 for k in range(K):
49 lo,hi=2**k, min(2**(k+1),L)
50 # clipped ranges, with exact valid-count normalization per position
51 idx=torch.arange(L,device=x.device)
52 a=(idx-hi).clamp(min=0); b=(idx-lo).clamp(min=0)
53 left=(pref[:,b+1]-pref[:,a+1]) # positions [i-hi+1,i-lo], then remove invalid below zero
54 left_count=(b-a).clamp(min=1).float()[None,:,None]
55 # Correct range is j in [i-hi, i-lo], exclusive invalid j<0.
56 left=(pref[:,(idx-lo+1).clamp(min=0)]-pref[:,(idx-hi).clamp(min=0)])/left_count
57 a=(idx+lo).clamp(max=L); b=(idx+hi).clamp(max=L)
58 right=(pref[:,b]-pref[:,a])/(b-a).clamp(min=1).float()[None,:,None]
59 ak=1/((k+1)*(k+2)*math.log(2))
60 out += ak*(left+right)
61 return out
62
63def local_mix(x,r=8):
64 B,L,D=x.shape; out=torch.zeros_like(x)
65 for d in range(1,r+1):
66 out[:,d:]+=x[:,:L-d]; out[:,:L-d]+=x[:,d:]
67 return out/(2*r)
68
69class Dyadic(nn.Module):
70 def __init__(self): super().__init__(); self.scale=nn.Parameter(torch.tensor(.5))
71 def forward(self,x): return x+self.scale*dyadic_mix(x)
72class Local(nn.Module):
73 def __init__(self,r=8): super().__init__(); self.scale=nn.Parameter(torch.tensor(.5)); self.r=r
74 def forward(self,x): return x+self.scale*local_mix(x,self.r)
75
76def train(model, L, steps=180, batch=64):
77 model.to(device); opt=torch.optim.Adam(model.parameters(),lr=.03); w=make_weights(L)
78 losses=[]; t=time.time()
79 for s in range(steps):
80 x=torch.randn(batch,L,1,device=device); y=long_target(x,w)
81 pred=model(x); loss=F.mse_loss(pred,y); opt.zero_grad(); loss.backward(); opt.step()
82 losses.append(float(loss))
83 with torch.no_grad():
84 x=torch.randn(512,L,1,device=device); y=long_target(x,w); val=float(F.mse_loss(model(x),y))
85 return {'final_train':losses[-1],'val_mse':val,'seconds':time.time()-t,'learned_scale':float(model.scale.detach().cpu())}
86
87def main():
88 out={'device':str(device),'math':math_check(),'experiments':{}}
89 for L in [64,256]:
90 # same architecture/optimization; local control has same single scalar parameter
91 out['experiments'][str(L)]={'local_radius8':train(Local(8),L),'dyadic_tail':train(Dyadic(),L)}
92 print(json.dumps(out,indent=2))
93if __name__=='__main__': main()